{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":56537,"databundleVersionId":8877088,"sourceType":"competition"},{"sourceId":8352923,"sourceType":"datasetVersion","datasetId":4866555},{"sourceId":8404301,"sourceType":"datasetVersion","datasetId":4998928},{"sourceId":8409068,"sourceType":"datasetVersion","datasetId":5004471},{"sourceId":8425368,"sourceType":"datasetVersion","datasetId":4865495},{"sourceId":8697069,"sourceType":"datasetVersion","datasetId":5171141},{"sourceId":8814346,"sourceType":"datasetVersion","datasetId":5302121},{"sourceId":175596903,"sourceType":"kernelVersion"}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# LEAP - Atmospheric Physics using AI\nClimate models are essential to understanding Earth’s climate system. This competition is to build machine learning models to emulate how the atmosphere works in climate models. This will help reduce the errors in climate trends and make more reliable future climate predictions. \n\nThis notebook aims to explore different approaches as baseline and understanding the competitions. The data in this competition is very large. \n\n### Table of Contents\n- [Giba Baseline XGBoost](https://www.kaggle.com/code/titericz/giba-baseline-xgboost)\n- [Neural Networks using Pytorch ](#simple_nn)\n\n\n### References\n- @AMEDEO BIOLATTI [Keras Baseline Seq2Seq](https://www.kaggle.com/code/abiolatti/keras-baseline-seq2seq)\n- @FARUKCAN SAGLAM [LEAP PyTorch Baseline](https://www.kaggle.com/code/greysky/leap-pytorch-baseline)\n- @LONNIE [LEAP Catboost Baseline](https://www.kaggle.com/code/lonnieqin/leap-catboost-baseline)\n- @EGOR TRUSHIN [[LEAP] FFNN/PyTorch](https://www.kaggle.com/code/egortrushin/leap-ffnn-pytorch)\n- @ANDONI IRAZUSTA [Pytorch NN](https://www.kaggle.com/code/airazusta014/pytorch-nn/notebook)\n- @GIBA [Giba Baseline XGBoost](https://www.kaggle.com/code/titericz/giba-baseline-xgboost)\n","metadata":{}},{"cell_type":"code","source":"import random, sys, gc, warnings, math, torch, os\nimport numpy as np\nimport pandas as pd\nimport polars as pl\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom sklearn import metrics\nimport tensorflow as tf\nimport keras\nimport warnings\nwarnings.simplefilter(action='ignore', category=FutureWarning)\npd.options.display.max_columns = None\n\n# Seed the same seed to all \ndef seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    keras.utils.set_random_seed(SEED)\n    tf.random.set_seed(SEED)\n    tf.config.experimental.enable_op_determinism()\n\nSEED = 42\nseed_everything(SEED)","metadata":{"execution":{"iopub.status.busy":"2024-06-29T23:58:20.137129Z","iopub.execute_input":"2024-06-29T23:58:20.137494Z","iopub.status.idle":"2024-06-29T23:58:26.504416Z","shell.execute_reply.started":"2024-06-29T23:58:20.137464Z","shell.execute_reply":"2024-06-29T23:58:26.503556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import ctypes\nfrom numba import cuda\nlibc = ctypes.CDLL(\"libc.so.6\")\ndef clear_memory():\n    libc.malloc_trim(0)\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-06-29T23:58:26.506279Z","iopub.execute_input":"2024-06-29T23:58:26.50684Z","iopub.status.idle":"2024-06-29T23:58:26.719106Z","shell.execute_reply.started":"2024-06-29T23:58:26.50681Z","shell.execute_reply":"2024-06-29T23:58:26.718128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Keras Seq2Seq\nThis approach trains Keras Seq2Seq model. Keras Seq2Seq (Sequence to Sequence) models are a type of neural network architecture used for tasks where input sequences are mapped to output sequences. ","metadata":{}},{"cell_type":"code","source":"import jax\nos.environ[\"KERAS_BACKEND\"] = \"jax\"        # Set jax backend\nprint(tf.__version__)\nprint(jax.__version__)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-29T23:58:26.720147Z","iopub.execute_input":"2024-06-29T23:58:26.720422Z","iopub.status.idle":"2024-06-29T23:58:26.725523Z","shell.execute_reply.started":"2024-06-29T23:58:26.720397Z","shell.execute_reply":"2024-06-29T23:58:26.724639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_files = [os.path.join(\"/kaggle/input/leap-train-tfrecords\", \"train_%.3d.tfrec\" % i) for i in range(100)]  # training dataset\nvalid_files = [os.path.join(\"/kaggle/input/leap-train-tfrecords\", \"train_%.3d.tfrec\" % i) for i in range(100, 101)]  # valid datasets","metadata":{"execution":{"iopub.status.busy":"2024-06-29T23:58:26.728149Z","iopub.execute_input":"2024-06-29T23:58:26.728502Z","iopub.status.idle":"2024-06-29T23:58:26.737013Z","shell.execute_reply.started":"2024-06-29T23:58:26.728467Z","shell.execute_reply":"2024-06-29T23:58:26.736203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    sample = pl.read_csv(os.path.join(\"/kaggle/input/leap-atmospheric-physics-ai-climsim/sample_submission.csv\"), n_rows=1)\n    TARGETS = sample.select(pl.exclude('sample_id')).columns\n    print(len(TARGETS))\n    BATCH_SIZE = 4096\n    learning_rate = 1e-3\n    epochs = 12   # epochs = 5\n    epochs_warmup = 1\n    epochs_ending = 2\n    steps_per_epoch = int(np.ceil(len(train_files) * 100_000 / BATCH_SIZE))\n\n    lr_scheduler = keras.optimizers.schedules.CosineDecay(\n        1e-4, # Initial learning rate\n        (epochs - epochs_warmup - epochs_ending) * steps_per_epoch, \n        warmup_target=learning_rate, # Warmup learning rate\n        warmup_steps=steps_per_epoch * epochs_warmup,\n        alpha=0.1\n    )","metadata":{"execution":{"iopub.status.busy":"2024-06-29T23:58:26.738121Z","iopub.execute_input":"2024-06-29T23:58:26.738398Z","iopub.status.idle":"2024-06-29T23:58:26.765943Z","shell.execute_reply.started":"2024-06-29T23:58:26.738373Z","shell.execute_reply":"2024-06-29T23:58:26.764983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\ndef parse_function(example_proto):\n    feature_description = {\n        'x': tf.io.FixedLenFeature([556], tf.float32),\n        'targets': tf.io.FixedLenFeature([368], tf.float32)\n    }\n    e = tf.io.parse_single_example(example_proto, feature_description)\n    return e['x'], e['targets']\n\ndef load_data(): \n    # Training dataset    \n    train_data_options = tf.data.Options()\n    train_data_options.deterministic = True\n\n    ds_train = (\n        tf.data.Dataset.from_tensor_slices(train_files)\n        .with_options(train_data_options)\n        .shuffle(100)\n        .interleave(\n            lambda file: tf.data.TFRecordDataset(file).map(parse_function,\n                                                           num_parallel_calls=tf.data.AUTOTUNE),\n            num_parallel_calls=tf.data.AUTOTUNE,\n            cycle_length=10,\n            block_length=1000,\n            deterministic=True\n        )\n        .shuffle(4 * CFG.BATCH_SIZE)\n        .batch(CFG.BATCH_SIZE)\n        .prefetch(tf.data.AUTOTUNE)\n    )\n    \n    # Valid dataset    \n    ds_valid = (\n        tf.data.TFRecordDataset(valid_files)\n        .map(parse_function)\n        .batch(CFG.BATCH_SIZE)\n        .prefetch(tf.data.AUTOTUNE)\n    )\n   \n    return ds_train, ds_valid","metadata":{"execution":{"iopub.status.busy":"2024-06-29T23:58:26.767439Z","iopub.execute_input":"2024-06-29T23:58:26.768167Z","iopub.status.idle":"2024-06-29T23:58:26.777099Z","shell.execute_reply.started":"2024-06-29T23:58:26.768128Z","shell.execute_reply":"2024-06-29T23:58:26.776099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train, ds_valid = load_data()\n\n# Normalize the training dataset's feature values\nnorm_x = keras.layers.Normalization()\nnorm_x.adapt(ds_train.map(lambda x, y: x).take(1000))","metadata":{"execution":{"iopub.status.busy":"2024-06-29T23:58:26.778299Z","iopub.execute_input":"2024-06-29T23:58:26.778586Z","iopub.status.idle":"2024-06-29T23:59:47.928229Z","shell.execute_reply.started":"2024-06-29T23:58:26.778561Z","shell.execute_reply":"2024-06-29T23:59:47.926988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" # normalize valid dataset's target values\nnorm_y = keras.layers.Normalization()\nnorm_y.adapt(ds_train.map(lambda x, y: y).take(1000))\n\nmean_y = norm_y.mean\nstdd_y = keras.ops.maximum(1e-10, norm_y.variance ** 0.5)","metadata":{"execution":{"iopub.status.busy":"2024-06-29T23:59:47.930257Z","iopub.execute_input":"2024-06-29T23:59:47.930708Z","iopub.status.idle":"2024-06-30T00:01:04.272032Z","shell.execute_reply.started":"2024-06-29T23:59:47.93067Z","shell.execute_reply":"2024-06-30T00:01:04.270628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model definition & Training","metadata":{}},{"cell_type":"code","source":"@keras.saving.register_keras_serializable(package=\"MyMetrics\", name=\"ClippedR2Score\")\nclass ClippedR2Score(keras.metrics.Metric):\n    def __init__(self, name='r2_score', **kwargs):\n        super().__init__(name=name, **kwargs)\n        self.base_metric = keras.metrics.R2Score(class_aggregation=None)\n        \n    def update_state(self, y_true, y_pred, sample_weight=None):\n        self.base_metric.update_state(y_true, y_pred, sample_weight=None)\n        \n    def result(self):\n        return keras.ops.mean(keras.ops.clip(self.base_metric.result(), 0.0, 1.0))\n        \n    def reset_states(self):\n        self.base_metric.reset_states()","metadata":{"execution":{"iopub.status.busy":"2024-06-30T00:01:04.273981Z","iopub.execute_input":"2024-06-30T00:01:04.274387Z","iopub.status.idle":"2024-06-30T00:01:04.28693Z","shell.execute_reply.started":"2024-06-30T00:01:04.27435Z","shell.execute_reply":"2024-06-30T00:01:04.285678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def x_to_seq(x):\n    x_seq0 = keras.ops.transpose(keras.ops.reshape(x[:, 0:60 * 6], (-1, 6, 60)), (0, 2, 1))\n    x_seq1 = keras.ops.transpose(keras.ops.reshape(x[:, 60 * 6 + 16:60 * 9 + 16], (-1, 3, 60)), (0, 2, 1))\n    x_flat = keras.ops.reshape(x[:, 60 * 6:60 * 6 + 16], (-1, 1, 16))\n    x_flat = keras.ops.repeat(x_flat, 60, axis=1)\n    return keras.ops.concatenate([x_seq0, x_seq1, x_flat], axis=-1)\n\ndef build_cnn(activation='relu'):    \n    return keras.Sequential([\n        keras.layers.Conv1D(256, 3, padding='same', activation=activation),\n        keras.layers.BatchNormalization(),\n        keras.layers.Conv1D(128, 3, padding='same', activation=activation),\n        keras.layers.BatchNormalization(),\n        keras.layers.Conv1D(64, 3, padding='same', activation=activation),\n        keras.layers.BatchNormalization(),\n    ])\n\ndef create_CNN_layers():\n    # Apply normalized X to training dataset\n    X_input = x = keras.layers.Input(ds_train.element_spec[0].shape[1:])\n    x = keras.layers.Normalization(mean=norm_x.mean, variance=norm_x.variance)(x)\n    x = x_to_seq(x)\n    \n    # Create CNN layers (64 filters)\n    e = e0 = keras.layers.Conv1D(64, 1, padding='same')(x)\n    e = build_cnn()(e)\n    # add global average to allow some comunication between all levels even in a small CNN\n    e = e0 + e + keras.layers.GlobalAveragePooling1D(keepdims=True)(e)\n    e = keras.layers.BatchNormalization()(e)\n    e = e + build_cnn()(e)\n    # Create another CNN layers (64 filters)\n    p_all = keras.layers.Conv1D(14, 1, padding='same')(e)\n\n    p_seq = p_all[:, :, :6]\n    p_seq = keras.ops.transpose(p_seq, (0, 2, 1))\n    p_seq = keras.layers.Flatten()(p_seq)\n    assert p_seq.shape[-1] == 360\n\n    p_flat = p_all[:, :, 6:6 + 8]\n    p_flat = keras.ops.mean(p_flat, axis=1)\n    assert p_flat.shape[-1] == 8\n\n    P = keras.ops.concatenate([p_seq, p_flat], axis=1)\n    return X_input, P\n","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-06-30T00:01:04.290973Z","iopub.execute_input":"2024-06-30T00:01:04.29232Z","iopub.status.idle":"2024-06-30T00:01:04.311074Z","shell.execute_reply.started":"2024-06-30T00:01:04.292277Z","shell.execute_reply":"2024-06-30T00:01:04.309823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"keras.utils.clear_session()\nX_input, P = create_CNN_layers()","metadata":{"execution":{"iopub.status.busy":"2024-06-30T00:01:04.312497Z","iopub.execute_input":"2024-06-30T00:01:04.31308Z","iopub.status.idle":"2024-06-30T00:01:04.760825Z","shell.execute_reply.started":"2024-06-30T00:01:04.313029Z","shell.execute_reply":"2024-06-30T00:01:04.759985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# build & compile\nmodel = keras.Model(X_input, P)\nmodel.compile(\n    loss='mse', \n    optimizer=keras.optimizers.Adam(CFG.lr_scheduler),\n    metrics=[ClippedR2Score()]\n)\nmodel.build(tuple(ds_train.element_spec[0].shape))\nmodel.summary()","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-06-30T00:01:04.761972Z","iopub.execute_input":"2024-06-30T00:01:04.76236Z","iopub.status.idle":"2024-06-30T00:01:04.829812Z","shell.execute_reply.started":"2024-06-30T00:01:04.762325Z","shell.execute_reply":"2024-06-30T00:01:04.828877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train_target_normalized = ds_train.map(lambda x, y: (x, (y - mean_y) / stdd_y))\nds_valid_target_normalized = ds_valid.map(lambda x, y: (x, (y - mean_y) / stdd_y))","metadata":{"execution":{"iopub.status.busy":"2024-06-30T00:01:04.831027Z","iopub.execute_input":"2024-06-30T00:01:04.831338Z","iopub.status.idle":"2024-06-30T00:01:04.88968Z","shell.execute_reply.started":"2024-06-30T00:01:04.83131Z","shell.execute_reply":"2024-06-30T00:01:04.888843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IS_TRAIN = True\n\nif IS_TRAIN:\n    history = model.fit(\n        ds_train_target_normalized,\n        validation_data=ds_valid_target_normalized,\n        epochs=CFG.epochs,\n        verbose=2,\n        callbacks=[\n            keras.callbacks.ModelCheckpoint(filepath='model.keras')\n        ]\n    )\nelse:\n    # Loads the weights\n    model.load_weights('/kaggle/input/leap-models/model.keras')","metadata":{"execution":{"iopub.status.busy":"2024-06-30T00:01:28.253128Z","iopub.execute_input":"2024-06-30T00:01:28.253804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_valid = np.concatenate([yb for _, yb in ds_valid])\np_valid = model.predict(ds_valid, batch_size=CFG.BATCH_SIZE) * stdd_y + mean_y","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"scores_valid = np.array([metrics.r2_score(y_valid[:, i], p_valid[:, i]) for i in range(len(CFG.TARGETS))])\nplt.plot(scores_valid.clip(-1, 1))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del y_valid, p_valid, ds_train, ds_valid\ngc.collect();","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Submission","metadata":{}},{"cell_type":"code","source":"df_test = (\n    pl.scan_csv(\"/kaggle/input/leap-atmospheric-physics-ai-climsim/test.csv\")\n    .select(pl.exclude(\"sample_id\"))\n    .cast(pl.Float32)\n    .collect()\n)  # Load test dataset","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Predict results\nmask = scores_valid <= 1e-3\np_test = model.predict(df_test.to_numpy(), batch_size=4 * CFG.BATCH_SIZE) * stdd_y + mean_y\np_test = np.array(p_test)\np_test[:, mask] = mean_y.numpy()[:, mask]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# correction of ptend_q0002 targets (from 12 to 29)\ndf_p_test = pd.DataFrame(p_test, columns=CFG.TARGETS)\n\nfor idx in range(12, 30):\n    df_p_test[f\"ptend_q0002_{idx}\"] = -df_test[f\"state_q0002_{idx}\"].to_numpy() / 1200\n    \np_test = df_p_test.values","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample = pl.read_csv(\"/kaggle/input/leap-atmospheric-physics-ai-climsim/sample_submission.csv\")\nsubmission = sample.to_pandas()\nsubmission[CFG.TARGETS] = submission[CFG.TARGETS] * p_test\npl.from_pandas(submission[[\"sample_id\"] + CFG.TARGETS]).write_csv(\"submission.csv\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# XGBoost/Catboost\nThis approach trains XGBoost model. ","metadata":{}},{"cell_type":"code","source":"import cudf, pickle\nfrom glob import glob\n\nimport xgboost as xgb\nfrom sklearn.metrics import r2_score\nfrom catboost import CatBoostRegressor","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load data","metadata":{}},{"cell_type":"code","source":"# Load test data    \ntest_files = glob(\"/kaggle/input/leap-dataset-giba/test_batch/*.parquet\")\ntest_df = pd.read_parquet(test_files[0]).astype('float32')\ntest_df = cudf.from_pandas(test_df) # Send to GPU for speedup\ndel test_files\ngc.collect()\n\n# Sample dataset contains all target variables\nsample_df = pd.read_csv(\"/kaggle/input/leap-atmospheric-physics-ai-climsim/sample_submission.csv\",\n                      nrows=1)\nsample_df.drop(['sample_id'],axis=1) # Drop sample_id column\nprint(f\"test_df.shape = {test_df.shape} sample.shape = {sample_df.shape}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    target_unpredictable = (\n        ['ptend_q0001_0','ptend_q0001_1','ptend_q0001_2','ptend_q0001_3','ptend_q0001_4','ptend_q0001_5','ptend_q0001_6','ptend_q0001_7','ptend_q0001_8','ptend_q0001_9','ptend_q0001_10','ptend_q0001_11'] +\n        ['ptend_q0002_0','ptend_q0002_1','ptend_q0002_2','ptend_q0002_3','ptend_q0002_4','ptend_q0002_5','ptend_q0002_6','ptend_q0002_7','ptend_q0002_8','ptend_q0002_9','ptend_q0002_10','ptend_q0002_11','ptend_q0002_12','ptend_q0002_13','ptend_q0002_14','ptend_q0002_15','ptend_q0002_16','ptend_q0002_17','ptend_q0002_18','ptend_q0002_19','ptend_q0002_20','ptend_q0002_21','ptend_q0002_22','ptend_q0002_23','ptend_q0002_24','ptend_q0002_25','ptend_q0002_26'] +\n        ['ptend_q0003_0','ptend_q0003_1','ptend_q0003_2','ptend_q0003_3','ptend_q0003_4','ptend_q0003_5','ptend_q0003_6','ptend_q0003_7','ptend_q0003_8','ptend_q0003_9','ptend_q0003_10','ptend_q0003_11'] +\n        ['ptend_u_0','ptend_u_1','ptend_u_2','ptend_u_3','ptend_u_4','ptend_u_5','ptend_u_6','ptend_u_7','ptend_u_8','ptend_u_9','ptend_u_10','ptend_u_11'] +\n        ['ptend_v_0','ptend_v_1','ptend_v_2','ptend_v_3','ptend_v_4','ptend_v_5','ptend_v_6','ptend_v_7','ptend_v_8','ptend_v_9','ptend_v_10','ptend_v_11']\n    )\n    print(target_unpredictable)\n    xgb_params = {\n                    'n_estimators': 200, \n                    'learning_rate': 0.1,\n                    'max_depth': 8,\n                    'device': 'cuda',\n                    'subsample': 0.40,\n                    'colsample_bytree': 0.95,\n                }\n    features = list(test_df.drop(['sample_id'],axis=1).columns)\n    all_targets = list(sample_df.drop(['sample_id'],axis=1).columns)\n\n# Take out unpredictable targets from all targets\nCFG.targets = [target for target in CFG.all_targets if target not in CFG.target_unpredictable]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Total number of features:\", len(CFG.features))\nprint(\"Total number of targets:\", len(CFG.targets))\nprint(\"Total number of all targets:\", len(CFG.all_targets))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model training and inference","metadata":{}},{"cell_type":"code","source":"class ModelTrainer:\n    def __init__(self, train_df, valid_df, test_df):\n        self.train_df = train_df\n        self.valid_df = valid_df\n        self.test_df = test_df\n        self.partial_targets = []\n    \n    # Train the model one target at each time\n    def train_model_for_one_target(self, target):\n        feature = target + '_pred'  # Feature name \n        target_mean = self.train_df[target].mean()\n\n        # Create the XGBoost regressor model\n        model = xgb.XGBRegressor(**CFG.xgb_params,\n                                 objective='reg:squarederror', \n                                 callbacks=[xgb.callback.EarlyStopping(rounds=30,\n                                                                       min_delta=1e-3,\n                                                                       save_best=False,\n                                                                       maximize=False)] # Early stopping\n                                )\n#         model = CatBoostRegressor(**CFG.cat_params)\n        # Train the model with one target \n        model.fit(\n            self.train_df[CFG.features],\n            self.train_df[target],\n            eval_set=[(self.valid_df[CFG.features], self.valid_df[target])],\n            verbose=False,\n        )\n\n        # model inference\n        self.valid_df[feature] = model.predict(self.valid_df[CFG.features])\n        valid_score = r2_score(self.valid_df[target].values_host, # values_host: numpy array\n                               self.valid_df[feature].values_host, force_finite=True)\n\n        # If r2 is postive, just use mean(target)\n        if valid_score > 0:\n            self.test_df[target] = model.predict(self.test_df[CFG.features])\n        else:\n            self.valid_df[feature] = 0. # Set to zero\n            self.test_df[target] = 0.\n\n        best_iter = model.best_iteration\n        del model; gc.collect()\n        self.partial_targets.append(target)\n        # Evaluate the model's predictions with r2 score\n        target_score = r2_score(self.valid_df[target].values_host, \n                                self.valid_df[feature].values_host, force_finite=True)\n        \n        # Overall score across all targets\n        overall_score = r2_score(self.valid_df[self.partial_targets].values_host, \n                                 self.valid_df[[target + '_pred' for target in self.partial_targets]].values_host)\n        print(f\"target = {target} r2(overall): {overall_score:.4f} / r2({target}): {target_score:.4f}\"\n              f\" best iter: {best_iter}\")\n        \n    def train_model(self):\n        for tnum, target in enumerate(CFG.targets):\n            # Model training and predictions for a single target column\n            if target not in CFG.target_unpredictable:\n                self.train_model_for_one_target(target)\n            else: # Target unpredictable\n                self.valid_df[target] = 0.\n                self.valid_df[target + '_pred'] = 0.\n                self.test_df[target] = 0.\n                \n        # Evaluate the model for all target columns\n        valid_score = r2_score(self.valid_df[CFG.targets].values_host, \n                              self.valid_df[[target + '_pred' for target in CFG.targets]].values_host)\n        print(f\"valid_score = {valid_score}\")\n        # Write out the valid to file\n        self.valid_df.to_pandas().to_parquet(f'validation_{valid_score:.4f}.parquet')\n               ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IS_TRAIN = False\nif IS_TRAIN:\n    # Load training data\n    train_files = sorted(glob(\"/kaggle/input/leap-dataset-giba/train_batch/*.parquet\"))\n    # Train on 2/17 of the full dataset\n    train_df = pd.read_parquet(train_files[:1]).astype('float32')\n    train_df = cudf.from_pandas(train_df) # Send to GPU for speedup\n    gc.collect()\n\n    # Validate on last file (625000 samples)\n    valid_df = pd.read_parquet(train_files[-1]).astype('float32')\n    valid_df = cudf.from_pandas(valid_df) # Send to GPU for speedup\n    del train_files\n    gc.collect()\n    print(f\"train_df.shape = {train_df.shape}, valid_df.shape = {valid_df.shape}, \")\n\n    trainer = ModelTrainer(train_df, valid_df, test_df)\n    trainer.train_model()\n    del train_df, valid_df; clear_memory()\n    \n    # Write test_df\n    path = 'xgboost/test_df.parquet'\n    test_df.to_pandas().to_parquet(path)\n    print(f\"Save test_df to {path}\")\n    clear_memory()","metadata":{"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IS_SUBMISSION = False\n\nif IS_SUBMISSION:\n    path = 'xgboost/test_df.parquet'\n    test_df = pd.read_parquet(glob(path)).astype('float32')\n    \n    # Update the submission with test_df\n    weights = sample_df.T \n    #print(f\"weights = {weights}\")\n    weights = weights.to_dict()[0]\n    #print(f\"weights = {weights}\")\n    \n    # Create a submission\n    submission = pd.read_csv(\"/kaggle/input/leap-atmospheric-physics-ai-climsim/sample_submission.csv\")\n    for target in tqdm(CFG.all_targets):\n        if (target not in CFG.target_unpredictable) and (weights[target] > 0) :\n            submission[target] = test_df[target].values\n            print(f\"target = {target}, test_df[target].values[1:5] = {test_df[target].values[1:5]}\")\n        else:\n            submission[target] = 0.\n\n    submission.to_csv('submission.csv', index=False)        \n    display(submission.head())\n    print(\"Save submission\")","metadata":{"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del test_df, sample_df\nclear_memory()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Neural Networks using Pytorch <a class='anchor' id='simple_nn'></a>\nThis approach trains a feed-forward Neural Network using Pytorch library.","metadata":{}},{"cell_type":"markdown","source":"### Imports and Configs <a class='anchor' id='s_nn_imports'></a>","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom time import time\n\nfrom pathlib import Path\nimport polars as pl\n# Pytorch\nimport torch\nfrom torch import nn\nfrom torch import optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import random_split","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    DEVICE = torch.device(\"cuda\") if torch.cuda.is_available() else torch.device(\"cpu\")\n    print(f\"Device: {DEVICE}\")\n    \n    NUM_EPOCHS = 50           # number of epochs\n    BATCH_SIZE = 1024         # batch size\n    LEARNING_RATE = 1e-4\n    ERR = 1e-6\n    \n    # TARGET WEIGHTS\n#     target_weights = [30981.265271661872, 22502.432413914863, 18894.14713004499, 14514.244730542465, 10944.348069459196, 9065.01072024503, 9663.669038687454, 12688.557362943708, 19890.17226527665, 25831.37317235381, 33890.367561807274, 44122.94111025334, 59811.25595068309, 79434.07500078829, 107358.80916894016, 135720.8418348218, 149399.8411114814, 128492.95185325432, 91746.23687305572, 72748.76911097553, 66531.53596840335, 62932.30598423903, 56610.26874314136, 49473.14369220607, 43029.18495420936, 36912.67491908133, 31486.93117928144, 26898.072997215502, 23316.638282978325, 20459.73133196152, 18385.68309639014, 17111.405107656312, 16337.80991958771, 15857.759882318944, 15580.902485189716, 15497.59045982052, 15612.2556996736, 15797.88455410361, 15974.218740897895, 16130.395527176632, 16261.310866446129, 16371.892401608216, 16397.019695140876, 16325.463899570548, 16228.641108112768, 16191.809643436269, 16341.207925934068, 16645.711351490587, 17005.493716683693, 17430.29874509864, 17907.24023203076, 18431.55334008694, 19032.471309392287, 19701.355113141435, 20408.236605392685, 20967.20795006453, 21194.427318009974, 21088.521528526755, 19437.91555757985, 13677.902713248171, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 871528441401.8333, 1083221770553.0684, 147034752676.7702, 35556045575.13566, 35153369257.41337, 46086368691.51654, 24689305171.692936, 11343276593.440475, 5396624651.94418, 2449353007.641508, 1132225885.703891, 579547849.1340877, 330219246.7861086, 207613930.3131764, 144580292.27473342, 109933282.92266414, 88706603.092171, 73819777.54163922, 63615988.74519494, 57250262.292053565, 52976073.06761927, 49653169.17819005, 46544975.11484598, 43167606.9599748, 39724375.20499403, 36317177.25886468, 33057511.80930482, 29869089.497658804, 26982386.85583376, 24416235.17215712, 22273651.697369896, 20553426.04804544, 19216240.03357431, 18167694.44812838, 17501855.536957663, 17169938.630597908, 17005382.258644175, 16998475.26752617, 17082890.987979066, 17227982.77516062, 17445823.21630204, 17757404.421785507, 18346092.75160569, 19400573.66632694, 20506722.48296608, 22469648.380506545, 23432031.455169585, 26204163.40545158, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 3673829810926.31, 371405570725.2526, 14219163611.984406, 3001863018.1934915, 1432766589.9326108, 884599805.0283787, 560127980.1033351, 386052567.7087711, 287331851.051439, 222703657.59538063, 181069239.6264349, 154620864.3164144, 138093777.60284117, 126605828.89875436, 117967840.02553518, 111005814.39518328, 105186901.20678852, 100168133.0295481, 95568646.67416307, 91457433.39515457, 88871610.45308323, 88829796.26374224, 91398113.73291488, 96585131.67000748, 104507692.01463065, 115895119.998433, 131939701.08213414, 154492946.00677127, 183147918.17086875, 215151374.22324687, 247158314.6345976, 266792879.42215955, 279115128.29108113, 370541510.87006927, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 877670509694.7871, 1174826943136.8308, 1270605570069.038, 21727315470.5208, 3159456646.5437946, 1090653401.282219, 727967089.8459107, 384399548.9506704, 290787296.9451616, 232703218.45048887, 197467462.7577736, 174310890.8025987, 160536437.73297343, 153567098.77048483, 152120124.9453068, 153115566.6756177, 153955545.42558223, 153734675.21565756, 154798666.36905554, 163346213.58113608, 180013139.3707387, 200324358.8534948, 220754613.1646765, 241290935.478592, 262868932.2066308, 284448910.01847774, 305681084.4142859, 327605088.8575117, 350473296.7263526, 373964594.1196182, 398396925.8173239, 423528355.65716046, 450447055.544388, 478857006.4973163, 508200335.7126168, 537309657.5789208, 566854568.2904652, 594618842.9455439, 619715928.2391286, 641395460.8414665, 663290039.7810476, 689274894.631561, 718208866.3397261, 743951200.8024124, 761776104.2945968, 772911224.3082078, 804001144.8046833, 772448774.7758856, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 4613823.568205323, 1999308.9343799097, 904636.2296014762, 433823.6123842511, 207201.39055371704, 107836.09164720173, 57647.915219220784, 40606.52305039815, 47739.86647922776, 51669.35493930698, 56438.19768395407, 60447.45665200092, 65251.4153955275, 71920.88588011517, 78529.58115204438, 83422.30217897324, 87036.98552475807, 90389.72631774022, 93982.39165674087, 97578.0099352472, 101428.21366062944, 104630.69200130588, 105685.04322626138, 103962.58423268417, 99650.31670632094, 94290.49986206587, 89514.90144353417, 85905.45713126978, 82784.9857650212, 79152.28707014346, 74847.81017353121, 70378.81859610273, 65420.04643792357, 59953.75184604176, 54764.28281143022, 50362.51288353384, 46212.571031725325, 41997.52779088816, 37692.05148110484, 33834.73460995647, 31846.09764364542, 31934.145655397457, 31454.81247448105, 30105.4073072481, 26957.830283611693, 27760.04479210889, 29853.374336459365, 19133.428743715107, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 7619940.584531054, 3148394.472742347, 1308415.0022178134, 540515.7720745018, 215237.1053603881, 102546.7276372816, 68453.67122640925, 50692.59053608593, 51487.52043139844, 52104.76838400132, 54019.39151917722, 55856.02168787862, 60347.30240270209, 68990.96019017675, 79096.88768563846, 87574.33453690328, 94158.56052476274, 101903.63670531697, 111746.9753834774, 122460.65399236557, 132086.69387474353, 141041.48571028374, 146354.09441287292, 145953.09590059065, 139496.8007888401, 128508.85108217449, 116665.51769667884, 107458.39706309135, 100259.97236694951, 94108.98505029618, 88439.89456238014, 82734.9027659809, 77061.08621371102, 71333.5319243128, 65999.72532130677, 61798.9972058361, 58237.356419617165, 54715.10266341248, 50825.84431702935, 46059.17688689915, 40740.26050401376, 36335.80228304863, 33981.57568605091, 33589.7143390849, 33988.88524112733, 36272.9364507092, 41183.34413717943, 29194.12369278645, 0.0040536134869726, 0.0138824238058072, 135129884.5084534, 12219717.5342461, 0.0090705273332672, 0.0085898851680217, 0.0215368188774867, 0.0336321308942602]\n\n    # Column names\n    feature_cols = ['state_t_0', 'state_t_1', 'state_t_2', 'state_t_3', 'state_t_4', 'state_t_5', 'state_t_6', 'state_t_7', 'state_t_8', 'state_t_9', 'state_t_10', 'state_t_11', 'state_t_12', 'state_t_13', 'state_t_14', 'state_t_15', 'state_t_16', 'state_t_17', 'state_t_18', 'state_t_19', 'state_t_20', 'state_t_21', 'state_t_22', 'state_t_23', 'state_t_24', 'state_t_25', 'state_t_26', 'state_t_27', 'state_t_28', 'state_t_29', 'state_t_30', 'state_t_31', 'state_t_32', 'state_t_33', 'state_t_34', 'state_t_35', 'state_t_36', 'state_t_37', 'state_t_38', 'state_t_39', 'state_t_40', 'state_t_41', 'state_t_42', 'state_t_43', 'state_t_44', 'state_t_45', 'state_t_46', 'state_t_47', 'state_t_48', 'state_t_49', 'state_t_50', 'state_t_51', 'state_t_52', 'state_t_53', 'state_t_54', 'state_t_55', 'state_t_56', 'state_t_57', 'state_t_58', 'state_t_59', 'state_q0001_0', 'state_q0001_1', 'state_q0001_2', 'state_q0001_3', 'state_q0001_4', 'state_q0001_5', 'state_q0001_6', 'state_q0001_7', 'state_q0001_8', 'state_q0001_9', 'state_q0001_10', 'state_q0001_11', 'state_q0001_12', 'state_q0001_13', 'state_q0001_14', 'state_q0001_15', 'state_q0001_16', 'state_q0001_17', 'state_q0001_18', 'state_q0001_19', 'state_q0001_20', 'state_q0001_21', 'state_q0001_22', 'state_q0001_23', 'state_q0001_24', 'state_q0001_25', 'state_q0001_26', 'state_q0001_27', 'state_q0001_28', 'state_q0001_29', 'state_q0001_30', 'state_q0001_31', 'state_q0001_32', 'state_q0001_33', 'state_q0001_34', 'state_q0001_35', 'state_q0001_36', 'state_q0001_37', 'state_q0001_38', 'state_q0001_39', 'state_q0001_40', 'state_q0001_41', 'state_q0001_42', 'state_q0001_43', 'state_q0001_44', 'state_q0001_45', 'state_q0001_46', 'state_q0001_47', 'state_q0001_48', 'state_q0001_49', 'state_q0001_50', 'state_q0001_51', 'state_q0001_52', 'state_q0001_53', 'state_q0001_54', 'state_q0001_55', 'state_q0001_56', 'state_q0001_57', 'state_q0001_58', 'state_q0001_59', 'state_q0002_0', 'state_q0002_1', 'state_q0002_2', 'state_q0002_3', 'state_q0002_4', 'state_q0002_5', 'state_q0002_6', 'state_q0002_7', 'state_q0002_8', 'state_q0002_9', 'state_q0002_10', 'state_q0002_11', 'state_q0002_12', 'state_q0002_13', 'state_q0002_14', 'state_q0002_15', 'state_q0002_16', 'state_q0002_17', 'state_q0002_18', 'state_q0002_19', 'state_q0002_20', 'state_q0002_21', 'state_q0002_22', 'state_q0002_23', 'state_q0002_24', 'state_q0002_25', 'state_q0002_26', 'state_q0002_27', 'state_q0002_28', 'state_q0002_29', 'state_q0002_30', 'state_q0002_31', 'state_q0002_32', 'state_q0002_33', 'state_q0002_34', 'state_q0002_35', 'state_q0002_36', 'state_q0002_37', 'state_q0002_38', 'state_q0002_39', 'state_q0002_40', 'state_q0002_41', 'state_q0002_42', 'state_q0002_43', 'state_q0002_44', 'state_q0002_45', 'state_q0002_46', 'state_q0002_47', 'state_q0002_48', 'state_q0002_49', 'state_q0002_50', 'state_q0002_51', 'state_q0002_52', 'state_q0002_53', 'state_q0002_54', 'state_q0002_55', 'state_q0002_56', 'state_q0002_57', 'state_q0002_58', 'state_q0002_59', 'state_q0003_0', 'state_q0003_1', 'state_q0003_2', 'state_q0003_3', 'state_q0003_4', 'state_q0003_5', 'state_q0003_6', 'state_q0003_7', 'state_q0003_8', 'state_q0003_9', 'state_q0003_10', 'state_q0003_11', 'state_q0003_12', 'state_q0003_13', 'state_q0003_14', 'state_q0003_15', 'state_q0003_16', 'state_q0003_17', 'state_q0003_18', 'state_q0003_19', 'state_q0003_20', 'state_q0003_21', 'state_q0003_22', 'state_q0003_23', 'state_q0003_24', 'state_q0003_25', 'state_q0003_26', 'state_q0003_27', 'state_q0003_28', 'state_q0003_29', 'state_q0003_30', 'state_q0003_31', 'state_q0003_32', 'state_q0003_33', 'state_q0003_34', 'state_q0003_35', 'state_q0003_36', 'state_q0003_37', 'state_q0003_38', 'state_q0003_39', 'state_q0003_40', 'state_q0003_41', 'state_q0003_42', 'state_q0003_43', 'state_q0003_44', 'state_q0003_45', 'state_q0003_46', 'state_q0003_47', 'state_q0003_48', 'state_q0003_49', 'state_q0003_50', 'state_q0003_51', 'state_q0003_52', 'state_q0003_53', 'state_q0003_54', 'state_q0003_55', 'state_q0003_56', 'state_q0003_57', 'state_q0003_58', 'state_q0003_59', 'state_u_0', 'state_u_1', 'state_u_2', 'state_u_3', 'state_u_4', 'state_u_5', 'state_u_6', 'state_u_7', 'state_u_8', 'state_u_9', 'state_u_10', 'state_u_11', 'state_u_12', 'state_u_13', 'state_u_14', 'state_u_15', 'state_u_16', 'state_u_17', 'state_u_18', 'state_u_19', 'state_u_20', 'state_u_21', 'state_u_22', 'state_u_23', 'state_u_24', 'state_u_25', 'state_u_26', 'state_u_27', 'state_u_28', 'state_u_29', 'state_u_30', 'state_u_31', 'state_u_32', 'state_u_33', 'state_u_34', 'state_u_35', 'state_u_36', 'state_u_37', 'state_u_38', 'state_u_39', 'state_u_40', 'state_u_41', 'state_u_42', 'state_u_43', 'state_u_44', 'state_u_45', 'state_u_46', 'state_u_47', 'state_u_48', 'state_u_49', 'state_u_50', 'state_u_51', 'state_u_52', 'state_u_53', 'state_u_54', 'state_u_55', 'state_u_56', 'state_u_57', 'state_u_58', 'state_u_59', 'state_v_0', 'state_v_1', 'state_v_2', 'state_v_3', 'state_v_4', 'state_v_5', 'state_v_6', 'state_v_7', 'state_v_8', 'state_v_9', 'state_v_10', 'state_v_11', 'state_v_12', 'state_v_13', 'state_v_14', 'state_v_15', 'state_v_16', 'state_v_17', 'state_v_18', 'state_v_19', 'state_v_20', 'state_v_21', 'state_v_22', 'state_v_23', 'state_v_24', 'state_v_25', 'state_v_26', 'state_v_27', 'state_v_28', 'state_v_29', 'state_v_30', 'state_v_31', 'state_v_32', 'state_v_33', 'state_v_34', 'state_v_35', 'state_v_36', 'state_v_37', 'state_v_38', 'state_v_39', 'state_v_40', 'state_v_41', 'state_v_42', 'state_v_43', 'state_v_44', 'state_v_45', 'state_v_46', 'state_v_47', 'state_v_48', 'state_v_49', 'state_v_50', 'state_v_51', 'state_v_52', 'state_v_53', 'state_v_54', 'state_v_55', 'state_v_56', 'state_v_57', 'state_v_58', 'state_v_59', 'state_ps', 'pbuf_SOLIN', 'pbuf_LHFLX', 'pbuf_SHFLX', 'pbuf_TAUX', 'pbuf_TAUY', 'pbuf_COSZRS', 'cam_in_ALDIF', 'cam_in_ALDIR', 'cam_in_ASDIF', 'cam_in_ASDIR', 'cam_in_LWUP', 'cam_in_ICEFRAC', 'cam_in_LANDFRAC', 'cam_in_OCNFRAC', 'cam_in_SNOWHLAND', 'pbuf_ozone_0', 'pbuf_ozone_1', 'pbuf_ozone_2', 'pbuf_ozone_3', 'pbuf_ozone_4', 'pbuf_ozone_5', 'pbuf_ozone_6', 'pbuf_ozone_7', 'pbuf_ozone_8', 'pbuf_ozone_9', 'pbuf_ozone_10', 'pbuf_ozone_11', 'pbuf_ozone_12', 'pbuf_ozone_13', 'pbuf_ozone_14', 'pbuf_ozone_15', 'pbuf_ozone_16', 'pbuf_ozone_17', 'pbuf_ozone_18', 'pbuf_ozone_19', 'pbuf_ozone_20', 'pbuf_ozone_21', 'pbuf_ozone_22', 'pbuf_ozone_23', 'pbuf_ozone_24', 'pbuf_ozone_25', 'pbuf_ozone_26', 'pbuf_ozone_27', 'pbuf_ozone_28', 'pbuf_ozone_29', 'pbuf_ozone_30', 'pbuf_ozone_31', 'pbuf_ozone_32', 'pbuf_ozone_33', 'pbuf_ozone_34', 'pbuf_ozone_35', 'pbuf_ozone_36', 'pbuf_ozone_37', 'pbuf_ozone_38', 'pbuf_ozone_39', 'pbuf_ozone_40', 'pbuf_ozone_41', 'pbuf_ozone_42', 'pbuf_ozone_43', 'pbuf_ozone_44', 'pbuf_ozone_45', 'pbuf_ozone_46', 'pbuf_ozone_47', 'pbuf_ozone_48', 'pbuf_ozone_49', 'pbuf_ozone_50', 'pbuf_ozone_51', 'pbuf_ozone_52', 'pbuf_ozone_53', 'pbuf_ozone_54', 'pbuf_ozone_55', 'pbuf_ozone_56', 'pbuf_ozone_57', 'pbuf_ozone_58', 'pbuf_ozone_59', 'pbuf_CH4_0', 'pbuf_CH4_1', 'pbuf_CH4_2', 'pbuf_CH4_3', 'pbuf_CH4_4', 'pbuf_CH4_5', 'pbuf_CH4_6', 'pbuf_CH4_7', 'pbuf_CH4_8', 'pbuf_CH4_9', 'pbuf_CH4_10', 'pbuf_CH4_11', 'pbuf_CH4_12', 'pbuf_CH4_13', 'pbuf_CH4_14', 'pbuf_CH4_15', 'pbuf_CH4_16', 'pbuf_CH4_17', 'pbuf_CH4_18', 'pbuf_CH4_19', 'pbuf_CH4_20', 'pbuf_CH4_21', 'pbuf_CH4_22', 'pbuf_CH4_23', 'pbuf_CH4_24', 'pbuf_CH4_25', 'pbuf_CH4_26', 'pbuf_CH4_27', 'pbuf_CH4_28', 'pbuf_CH4_29', 'pbuf_CH4_30', 'pbuf_CH4_31', 'pbuf_CH4_32', 'pbuf_CH4_33', 'pbuf_CH4_34', 'pbuf_CH4_35', 'pbuf_CH4_36', 'pbuf_CH4_37', 'pbuf_CH4_38', 'pbuf_CH4_39', 'pbuf_CH4_40', 'pbuf_CH4_41', 'pbuf_CH4_42', 'pbuf_CH4_43', 'pbuf_CH4_44', 'pbuf_CH4_45', 'pbuf_CH4_46', 'pbuf_CH4_47', 'pbuf_CH4_48', 'pbuf_CH4_49', 'pbuf_CH4_50', 'pbuf_CH4_51', 'pbuf_CH4_52', 'pbuf_CH4_53', 'pbuf_CH4_54', 'pbuf_CH4_55', 'pbuf_CH4_56', 'pbuf_CH4_57', 'pbuf_CH4_58', 'pbuf_CH4_59', 'pbuf_N2O_0', 'pbuf_N2O_1', 'pbuf_N2O_2', 'pbuf_N2O_3', 'pbuf_N2O_4', 'pbuf_N2O_5', 'pbuf_N2O_6', 'pbuf_N2O_7', 'pbuf_N2O_8', 'pbuf_N2O_9', 'pbuf_N2O_10', 'pbuf_N2O_11', 'pbuf_N2O_12', 'pbuf_N2O_13', 'pbuf_N2O_14', 'pbuf_N2O_15', 'pbuf_N2O_16', 'pbuf_N2O_17', 'pbuf_N2O_18', 'pbuf_N2O_19', 'pbuf_N2O_20', 'pbuf_N2O_21', 'pbuf_N2O_22', 'pbuf_N2O_23', 'pbuf_N2O_24', 'pbuf_N2O_25', 'pbuf_N2O_26', 'pbuf_N2O_27', 'pbuf_N2O_28', 'pbuf_N2O_29', 'pbuf_N2O_30', 'pbuf_N2O_31', 'pbuf_N2O_32', 'pbuf_N2O_33', 'pbuf_N2O_34', 'pbuf_N2O_35', 'pbuf_N2O_36', 'pbuf_N2O_37', 'pbuf_N2O_38', 'pbuf_N2O_39', 'pbuf_N2O_40', 'pbuf_N2O_41', 'pbuf_N2O_42', 'pbuf_N2O_43', 'pbuf_N2O_44', 'pbuf_N2O_45', 'pbuf_N2O_46', 'pbuf_N2O_47', 'pbuf_N2O_48', 'pbuf_N2O_49', 'pbuf_N2O_50', 'pbuf_N2O_51', 'pbuf_N2O_52', 'pbuf_N2O_53', 'pbuf_N2O_54', 'pbuf_N2O_55', 'pbuf_N2O_56', 'pbuf_N2O_57', 'pbuf_N2O_58', 'pbuf_N2O_59']\n    target_cols = ['ptend_t_0', 'ptend_t_1', 'ptend_t_2', 'ptend_t_3', 'ptend_t_4', 'ptend_t_5', 'ptend_t_6', 'ptend_t_7', 'ptend_t_8', 'ptend_t_9', 'ptend_t_10', 'ptend_t_11', 'ptend_t_12', 'ptend_t_13', 'ptend_t_14', 'ptend_t_15', 'ptend_t_16', 'ptend_t_17', 'ptend_t_18', 'ptend_t_19', 'ptend_t_20', 'ptend_t_21', 'ptend_t_22', 'ptend_t_23', 'ptend_t_24', 'ptend_t_25', 'ptend_t_26', 'ptend_t_27', 'ptend_t_28', 'ptend_t_29', 'ptend_t_30', 'ptend_t_31', 'ptend_t_32', 'ptend_t_33', 'ptend_t_34', 'ptend_t_35', 'ptend_t_36', 'ptend_t_37', 'ptend_t_38', 'ptend_t_39', 'ptend_t_40', 'ptend_t_41', 'ptend_t_42', 'ptend_t_43', 'ptend_t_44', 'ptend_t_45', 'ptend_t_46', 'ptend_t_47', 'ptend_t_48', 'ptend_t_49', 'ptend_t_50', 'ptend_t_51', 'ptend_t_52', 'ptend_t_53', 'ptend_t_54', 'ptend_t_55', 'ptend_t_56', 'ptend_t_57', 'ptend_t_58', 'ptend_t_59', 'ptend_q0001_0', 'ptend_q0001_1', 'ptend_q0001_2', 'ptend_q0001_3', 'ptend_q0001_4', 'ptend_q0001_5', 'ptend_q0001_6', 'ptend_q0001_7', 'ptend_q0001_8', 'ptend_q0001_9', 'ptend_q0001_10', 'ptend_q0001_11', 'ptend_q0001_12', 'ptend_q0001_13', 'ptend_q0001_14', 'ptend_q0001_15', 'ptend_q0001_16', 'ptend_q0001_17', 'ptend_q0001_18', 'ptend_q0001_19', 'ptend_q0001_20', 'ptend_q0001_21', 'ptend_q0001_22', 'ptend_q0001_23', 'ptend_q0001_24', 'ptend_q0001_25', 'ptend_q0001_26', 'ptend_q0001_27', 'ptend_q0001_28', 'ptend_q0001_29', 'ptend_q0001_30', 'ptend_q0001_31', 'ptend_q0001_32', 'ptend_q0001_33', 'ptend_q0001_34', 'ptend_q0001_35', 'ptend_q0001_36', 'ptend_q0001_37', 'ptend_q0001_38', 'ptend_q0001_39', 'ptend_q0001_40', 'ptend_q0001_41', 'ptend_q0001_42', 'ptend_q0001_43', 'ptend_q0001_44', 'ptend_q0001_45', 'ptend_q0001_46', 'ptend_q0001_47', 'ptend_q0001_48', 'ptend_q0001_49', 'ptend_q0001_50', 'ptend_q0001_51', 'ptend_q0001_52', 'ptend_q0001_53', 'ptend_q0001_54', 'ptend_q0001_55', 'ptend_q0001_56', 'ptend_q0001_57', 'ptend_q0001_58', 'ptend_q0001_59', 'ptend_q0002_0', 'ptend_q0002_1', 'ptend_q0002_2', 'ptend_q0002_3', 'ptend_q0002_4', 'ptend_q0002_5', 'ptend_q0002_6', 'ptend_q0002_7', 'ptend_q0002_8', 'ptend_q0002_9', 'ptend_q0002_10', 'ptend_q0002_11', 'ptend_q0002_12', 'ptend_q0002_13', 'ptend_q0002_14', 'ptend_q0002_15', 'ptend_q0002_16', 'ptend_q0002_17', 'ptend_q0002_18', 'ptend_q0002_19', 'ptend_q0002_20', 'ptend_q0002_21', 'ptend_q0002_22', 'ptend_q0002_23', 'ptend_q0002_24', 'ptend_q0002_25', 'ptend_q0002_26', 'ptend_q0002_27', 'ptend_q0002_28', 'ptend_q0002_29', 'ptend_q0002_30', 'ptend_q0002_31', 'ptend_q0002_32', 'ptend_q0002_33', 'ptend_q0002_34', 'ptend_q0002_35', 'ptend_q0002_36', 'ptend_q0002_37', 'ptend_q0002_38', 'ptend_q0002_39', 'ptend_q0002_40', 'ptend_q0002_41', 'ptend_q0002_42', 'ptend_q0002_43', 'ptend_q0002_44', 'ptend_q0002_45', 'ptend_q0002_46', 'ptend_q0002_47', 'ptend_q0002_48', 'ptend_q0002_49', 'ptend_q0002_50', 'ptend_q0002_51', 'ptend_q0002_52', 'ptend_q0002_53', 'ptend_q0002_54', 'ptend_q0002_55', 'ptend_q0002_56', 'ptend_q0002_57', 'ptend_q0002_58', 'ptend_q0002_59', 'ptend_q0003_0', 'ptend_q0003_1', 'ptend_q0003_2', 'ptend_q0003_3', 'ptend_q0003_4', 'ptend_q0003_5', 'ptend_q0003_6', 'ptend_q0003_7', 'ptend_q0003_8', 'ptend_q0003_9', 'ptend_q0003_10', 'ptend_q0003_11', 'ptend_q0003_12', 'ptend_q0003_13', 'ptend_q0003_14', 'ptend_q0003_15', 'ptend_q0003_16', 'ptend_q0003_17', 'ptend_q0003_18', 'ptend_q0003_19', 'ptend_q0003_20', 'ptend_q0003_21', 'ptend_q0003_22', 'ptend_q0003_23', 'ptend_q0003_24', 'ptend_q0003_25', 'ptend_q0003_26', 'ptend_q0003_27', 'ptend_q0003_28', 'ptend_q0003_29', 'ptend_q0003_30', 'ptend_q0003_31', 'ptend_q0003_32', 'ptend_q0003_33', 'ptend_q0003_34', 'ptend_q0003_35', 'ptend_q0003_36', 'ptend_q0003_37', 'ptend_q0003_38', 'ptend_q0003_39', 'ptend_q0003_40', 'ptend_q0003_41', 'ptend_q0003_42', 'ptend_q0003_43', 'ptend_q0003_44', 'ptend_q0003_45', 'ptend_q0003_46', 'ptend_q0003_47', 'ptend_q0003_48', 'ptend_q0003_49', 'ptend_q0003_50', 'ptend_q0003_51', 'ptend_q0003_52', 'ptend_q0003_53', 'ptend_q0003_54', 'ptend_q0003_55', 'ptend_q0003_56', 'ptend_q0003_57', 'ptend_q0003_58', 'ptend_q0003_59', 'ptend_u_0', 'ptend_u_1', 'ptend_u_2', 'ptend_u_3', 'ptend_u_4', 'ptend_u_5', 'ptend_u_6', 'ptend_u_7', 'ptend_u_8', 'ptend_u_9', 'ptend_u_10', 'ptend_u_11', 'ptend_u_12', 'ptend_u_13', 'ptend_u_14', 'ptend_u_15', 'ptend_u_16', 'ptend_u_17', 'ptend_u_18', 'ptend_u_19', 'ptend_u_20', 'ptend_u_21', 'ptend_u_22', 'ptend_u_23', 'ptend_u_24', 'ptend_u_25', 'ptend_u_26', 'ptend_u_27', 'ptend_u_28', 'ptend_u_29', 'ptend_u_30', 'ptend_u_31', 'ptend_u_32', 'ptend_u_33', 'ptend_u_34', 'ptend_u_35', 'ptend_u_36', 'ptend_u_37', 'ptend_u_38', 'ptend_u_39', 'ptend_u_40', 'ptend_u_41', 'ptend_u_42', 'ptend_u_43', 'ptend_u_44', 'ptend_u_45', 'ptend_u_46', 'ptend_u_47', 'ptend_u_48', 'ptend_u_49', 'ptend_u_50', 'ptend_u_51', 'ptend_u_52', 'ptend_u_53', 'ptend_u_54', 'ptend_u_55', 'ptend_u_56', 'ptend_u_57', 'ptend_u_58', 'ptend_u_59', 'ptend_v_0', 'ptend_v_1', 'ptend_v_2', 'ptend_v_3', 'ptend_v_4', 'ptend_v_5', 'ptend_v_6', 'ptend_v_7', 'ptend_v_8', 'ptend_v_9', 'ptend_v_10', 'ptend_v_11', 'ptend_v_12', 'ptend_v_13', 'ptend_v_14', 'ptend_v_15', 'ptend_v_16', 'ptend_v_17', 'ptend_v_18', 'ptend_v_19', 'ptend_v_20', 'ptend_v_21', 'ptend_v_22', 'ptend_v_23', 'ptend_v_24', 'ptend_v_25', 'ptend_v_26', 'ptend_v_27', 'ptend_v_28', 'ptend_v_29', 'ptend_v_30', 'ptend_v_31', 'ptend_v_32', 'ptend_v_33', 'ptend_v_34', 'ptend_v_35', 'ptend_v_36', 'ptend_v_37', 'ptend_v_38', 'ptend_v_39', 'ptend_v_40', 'ptend_v_41', 'ptend_v_42', 'ptend_v_43', 'ptend_v_44', 'ptend_v_45', 'ptend_v_46', 'ptend_v_47', 'ptend_v_48', 'ptend_v_49', 'ptend_v_50', 'ptend_v_51', 'ptend_v_52', 'ptend_v_53', 'ptend_v_54', 'ptend_v_55', 'ptend_v_56', 'ptend_v_57', 'ptend_v_58', 'ptend_v_59', 'cam_out_NETSW', 'cam_out_FLWDS', 'cam_out_PRECSC', 'cam_out_PRECC', 'cam_out_SOLS', 'cam_out_SOLL', 'cam_out_SOLSD', 'cam_out_SOLLD']\n    \n    print(f\"len(feature_col) = {len(feature_cols)}, len(target_cols) = {len(target_cols)}\")\n     # Replace columns\n    replace_from = ['ptend_q0002_0', 'ptend_q0002_1', 'ptend_q0002_2', 'ptend_q0002_3', 'ptend_q0002_4', 'ptend_q0002_5', 'ptend_q0002_6', 'ptend_q0002_7', 'ptend_q0002_8', 'ptend_q0002_9', 'ptend_q0002_10', 'ptend_q0002_11', 'ptend_q0002_12', 'ptend_q0002_13', 'ptend_q0002_14', 'ptend_q0002_15', 'ptend_q0002_16', 'ptend_q0002_17', 'ptend_q0002_18', 'ptend_q0002_19', 'ptend_q0002_20', 'ptend_q0002_21', 'ptend_q0002_22', 'ptend_q0002_23', 'ptend_q0002_24', 'ptend_q0002_25', 'ptend_q0002_26']\n    replace_to = ['state_q0002_0', 'state_q0002_1', 'state_q0002_2', 'state_q0002_3', 'state_q0002_4', 'state_q0002_5', 'state_q0002_6', 'state_q0002_7', 'state_q0002_8', 'state_q0002_9', 'state_q0002_10', 'state_q0002_11', 'state_q0002_12', 'state_q0002_13', 'state_q0002_14', 'state_q0002_15', 'state_q0002_16', 'state_q0002_17', 'state_q0002_18', 'state_q0002_19', 'state_q0002_20', 'state_q0002_21', 'state_q0002_22', 'state_q0002_23', 'state_q0002_24', 'state_q0002_25', 'state_q0002_26']\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load training data <a class='anchor' id='load_data'></a>\nThe training dataset is very large, which can lead to longer loading times. The polars library is used to read the training dataset as a dataframe. \n\n\nThe first 557 columns relate to 25 input variables. The remaining 368 columns corresponding to 14 target variables. ","metadata":{}},{"cell_type":"code","source":"class LeapDataset(Dataset):\n    # y_weights: Weights to be applied to the target data.\n    def __init__(self, data):\n        \n        super().__init__()\n        \n        # Create X (Features)\n        self.x = data[CFG.feature_cols]\n        self.x = self.x.to_numpy()\n        self.x = torch.from_numpy(self.x)\n        # Create Y (Targets)\n        self.y = data[CFG.target_cols]\n        self.y = self.y.to_numpy()\n        self.y = torch.from_numpy(self.y)\n        # Scale Y with given weights. Y is tensor\n        # self.y = self.y * torch.tensor(CFG.target_weights)\n        \n        \n    def __getitem__(self, idx):\n        x = self.x[idx]\n        y = self.y[idx]\n        # Normalize X and Y\n        x = (x - self.x_mean) / self.x_std\n        y = (y - self.y_mean) / self.y_std\n        # Convert to 32 bit\n        x = x.to(torch.float32)\n        y = y.to(torch.float32)\n        \n        return x, y\n    \n    def __len__(self):\n        return len(self.y)\n    \n    def calc_mean_std(self):\n        # compute x mean\n        x_mean = torch.mean(self.x, 0)\n        x_std = torch.maximum(torch.std(self.x, 0), torch.tensor(CFG.ERR))\n        # Compute y_mean\n        y_mean = self.y.mean(axis=0)\n        y_std = torch.maximum(torch.sqrt(torch.mean(torch.pow(self.y, 2), 0)),\n                              torch.tensor(CFG.ERR))\n        print(\"Mean(x) Shape: \", x_mean.shape)\n        print(\"Mean(y) Shape: \", y_mean.shape)\n        print(\"Std(x) Shape: \", x_std.shape)\n        print(\"Std(y) Shape: \", y_std.shape)\n        self.x_mean, self.x_std, self.y_mean, self.y_std = x_mean, x_std, y_mean, y_std\n        return x_mean, x_std, y_mean, y_std","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_files = sorted(glob(\"/kaggle/input/leap-data-subste/Leap_data_chunk*/*.parquet\"))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load training\n# data = pd.read_parquet('/kaggle/input/leap-data-dataset/train_sampled_1m.parquet')\n# Load training data\ntrain_files = sorted(glob(\"/kaggle/input/leap-data-subste/Leap_data_chunk1.parquet*.parquet\"))\n    # Train on 2/17 of the full dataset\n    train_df = pd.read_parquet(train_files[:1]).astype('float32')\n\ndata = cudf.from_pandas(data) # Send to GPU for speedup\ngc.collect()\ndisplay(data.columns)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = LeapDataset(data) # Create a dataset","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_MEAN, X_STD, Y_MEAN, Y_STD = ds.calc_mean_std()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training the model","metadata":{}},{"cell_type":"code","source":"class LEAPModel(nn.Module): \n    def __init__(self, dims:list):\n        super().__init__()\n                \n        layers = []\n        for i in range(len(dims) - 2):\n            layers.append(nn.Linear(dims[i], dims[i + 1]))\n            layers.append(nn.LayerNorm(dims[i + 1]))\n            layers.append(nn.ReLU())\n            \n        layers.append(nn.Linear(dims[-2], dims[-1]))\n        self.network = nn.Sequential(*layers)\n\n    def forward(self, x):\n        return self.network(x)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LEAPTrainer:\n    def __init__(self, model):\n        self.model = model\n\n    #  Calculate the R^2 (coefficient of determination) regression score.\n    def r2_score(self, y_pred:torch.Tensor, y_true:torch.Tensor) -> float:\n        ss_res = torch.sum((y_true - y_pred) ** 2)\n        ss_tot = torch.sum((y_true - torch.mean(y_true)) ** 2)\n\n        r2 = 1 - ss_res / ss_tot\n\n        return r2.item() # R^2 score\n\n    # Train the deep learning model for 1 epoch.\n    def train_fn(self, train_loader: DataLoader, \n                 optimizer: optim.Optimizer, criterion: nn.Module, epoch: int):\n        progress_bar = tqdm(enumerate(train_loader, start=1), total=len(train_loader), ncols=100)\n        progress_bar.set_description(f'Epoch {epoch}')\n        self.model.train()\n        train_loss = 0\n        for step, batch in progress_bar:  # Process training data in batch\n            x, y = batch\n            x, y = x.to(CFG.DEVICE), y.to(CFG.DEVICE)  \n            # clear out the gradients of all Variables in this optimizer (i.e. W, b)\n            optimizer.zero_grad()\n            # Make predictions\n            y_pred = self.model(x)\n            loss = criterion(y_pred, y)\n            loss.backward()\n            optimizer.step()\n            print(f\"loss = {loss}\")\n            train_loss += loss.cpu().item()\n\n            progress_bar.set_postfix({\n                'train_loss': train_loss / step,\n            })\n\n        return train_loss\n\n    # Validate the deep learning model for 1 epoch.\n    def valid_fn(self, valid_loader: DataLoader, epoch:int) -> float:\n        progress_bar = tqdm(enumerate(valid_loader, start=1), total=len(valid_loader), ncols=100)\n        progress_bar.set_description(f'Epoch {epoch}')\n        self.model.eval()\n        val_score = 0\n        with torch.no_grad():\n            for step, batch in progress_bar:\n                x, y = batch\n                x, y = x.to(CFG.DEVICE), y.to(CFG.DEVICE)\n                # Make a predictions\n                y_pred = self.model(x)\n                # Scale 'y' with mean and std\n                y = y.cpu()\n                y = (y * Y_STD) + Y_MEAN\n\n                y_pred = y_pred.cpu()\n                y_pred[:, Y_STD < (1.1 * CFG.ERR)] = 0\n                y_pred = (y_pred * Y_STD) + Y_MEAN\n\n                val_score += self.r2_score(y_pred, y)\n\n                progress_bar.set_postfix({\n                    'valid_score': val_score / step,\n                })\n\n        return val_score\n\n    # Train a simple NN model \n    def train_model(self, ds: LeapDataset):\n        ds_train, ds_valid = random_split(ds, [0.8, 0.2]) # Split dataset into 80% training 20% valid dataset\n        print(f\"len(ds_train) = {len(ds_train)}, len(ds_valid) = {len(ds_valid)}\")\n        # Create data loader\n        train_loader = DataLoader(ds_train, batch_size=CFG.BATCH_SIZE, shuffle=True, drop_last=True)\n        valid_loader = DataLoader(ds_valid, batch_size=CFG.BATCH_SIZE, shuffle=False, drop_last=False)\n        # Select MSE loss as loss function, Adam optimizer, and polynomial scheduler\n        criterion = nn.MSELoss()\n        optimizer = optim.Adam(model.parameters(), lr=CFG.LEARNING_RATE)\n        scheduler = lr_scheduler.PolynomialLR(optimizer, power=1.0, total_iters=CFG.NUM_EPOCHS)\n\n        # Set initial best as infinity\n        best_score = -np.inf\n        # Training loop\n        for epoch in range(CFG.NUM_EPOCHS):\n            train_loss = self.train_fn(train_loader, optimizer, criterion, epoch)\n            val_score = self.valid_fn(valid_loader, epoch)\n            print(f\"epoch = {epoch}, train_loss = {train_loss}, val_score = {val_score}\")\n\n            if val_score > best_score:\n                best_score = val_score\n                # Save the best model state\n                model_path = f'best_nn_model.pth'\n                torch.save(model.state_dict(), model_path)\n                print(f\"Save the model to {model_path}\")\n                print(f\"--- best epoch = {epoch}, train_loss = {train_loss}, val_score = {val_score} ---\")\n            scheduler.step()\n        print(\"Complete training the model\")\n        del train_loader, valid_loader, ds_train, ds_valid\n    ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAINING = False\nif TRAINING:\n    model = LEAPModel([len(CFG.feature_cols), 1024, 512, len(CFG.target_cols)]) # hidden layer (1024x512)\n    model = model.to(CFG.DEVICE)\n    model.load_state_dict(torch.load(\"/kaggle/input/leap-feedforward-neural-network-baseline/best_model.pth\",\n                                     map_location=CFG.DEVICE))\n    trainer = LEAPTrainer(model)\n    trainer.train_model(ds)\n    os._exit(0)","metadata":{"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Inference ","metadata":{}},{"cell_type":"code","source":"class ModelInfer:\n    def __init__(self, model):\n        self.model = model\n        \n    # Generate the predictions\n    def pred_fn(self, test_df: pl.DataFrame):\n\n        self.model.eval()\n\n        x_test = test_df[CFG.feature_cols]\n        x_test = x_test.to_numpy()\n        x_test = torch.from_numpy(x_test)\n\n        x_test = (x_test - X_MEAN) / X_STD\n\n        x_test = x_test.to(torch.float32)\n        x_test = x_test.to(CFG.DEVICE)\n\n        with torch.no_grad():\n            y_pred = self.model(x_test)\n\n        y_pred = y_pred.cpu()\n        y_pred = y_pred.to(torch.float64)\n\n        y_pred[:, Y_STD < (1.1 * CFG.ERR)] = 0\n        y_pred = (y_pred * Y_STD) + Y_MEAN\n\n        y_pred = y_pred.detach()\n        y_pred = y_pred.cpu()\n        y_pred = y_pred.numpy()\n\n        return y_pred","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IS_SUBMISSION = False\n\nif IS_SUBMISSION:\n    # Load the model\n    model = LEAPModel([len(CFG.feature_cols), 1024, 512, len(CFG.target_cols)])\n    model = model.to(CFG.DEVICE)\n    model.load_state_dict(torch.load(\"/kaggle/input/leap-feedforward-neural-network-baseline/best_nn_model.pth\", map_location=CFG.DEVICE))\n    \n    # Load test dataset\n    test_df = pl.read_csv('/kaggle/input/leap-atmospheric-physics-ai-climsim/test.csv')\n    test_df = test_df.to_pandas() \n    test_df = test_df.set_index(\"sample_id\")\n    \n    # Make predictions\n    infer = ModelInfer(model)\n    preds = infer.pred_fn(test_df)\n    \n    # Load sample submission dataset\n    sample_df = pl.read_csv('/kaggle/input/leap-atmospheric-physics-ai-climsim/sample_submission.csv')\n    sample_df = sample_df.to_pandas()\n    sample_df = sample_df.set_index(\"sample_id\")\n    sample_df.loc[test_df.index, CFG.target_cols] = preds\n    \n    # Update the test_df with select columns\n#     static_pred = -test_df[CFG.replace_to].to_numpy() * sample_df[CFG.replace_from].to_numpy() / 1200\n    # Scale the submission results with magic scale\n    sample_df[CFG.replace_from] = static_pred\n    # Reset the id and columns\n    sample_df = sample_df.reset_index()\n    sample_df = sample_df[[\"sample_id\"] + CFG.target_cols]\n\n    submission_df = pl.from_pandas(sample_df)\n    submission_df.write_csv(\"submission.csv\")\n    display(submission_df)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}