{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.13"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":56537,"databundleVersionId":8877088,"sourceType":"competition"},{"sourceId":8409068,"sourceType":"datasetVersion","datasetId":5004471}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":29612.866972,"end_time":"2024-06-17T12:11:37.718808","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-06-17T03:58:04.851836","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# LEAP Climsim: Atmospheric Physics Prediction using Keras (JAX Backend)\n### Tree-based baselines with Ordinal Encoding and 5-Fold CV MAE evaluation\n\n## 1. Introduction  \nBrief overview of the goal — predicting atmospheric variables from high-dimensional input features using a Keras (JAX) model.\n\n## 2. Data Loading  \n- Mount or locate dataset paths\n- Load metadata or samples\n- Print dimensions and column names\n\n## 3. TFRecord Parsing & Dataset Preparation  \n- Define `_parse_function()`  \n- Build `train_files` / `valid_files` lists  \n- Create the `tf.data.Dataset` pipelines for train/valid  \n\n## 4. Model Architecture  \n- Define the Keras model (Sequential, Functional API, etc.)  \n- Print `model.summary()`  \n- Set optimizer, loss, and metrics  \n\n## 5. Model Training  \n- Configure callbacks (EarlyStopping, ReduceLROnPlateau)  \n- Train model with `model.fit()`  \n- Monitor loss curves  \n\n## 6. Evaluation  \n- Evaluate model performance on validation set  \n- Show sample predictions  \n- Compute regression metrics (MAE, RMSE, R², etc.)\n\n## 7. Submission  \n- Predict on test data  \n- Format submission file  \n- Save `.csv` for Kaggle submission  \n\n## 8. Conclusion  \nSummarize the key findings, performance, and potential improvements.  \nSuggest future experiments or model extensions.","metadata":{}},{"cell_type":"markdown","source":"# 1. Introduction\n\nUnderstanding and predicting atmospheric processes is a fundamental challenge in modern climate science. The LEAP Climsim dataset provides a large-scale benchmark for data-driven modeling of atmospheric physics using AI.\n\nIn this notebook, we develop a deep learning pipeline based on Keras (JAX backend) to predict complex multi-output atmospheric variables from high-dimensional input data.\n\nKey Features of This Notebook:\n\n* End-to-end TFRecord Pipeline: Efficient loading and preprocessing of the LEAP Climsim dataset using TensorFlow’s tf.data API.\n\n* Neural Network Model (Keras + JAX): A deterministic and optimized model setup leveraging JAX for numerical stability and GPU acceleration.\n\n* Multi-Target Regression: Simultaneous prediction of 368 atmospheric variables from 556 features.\n\n* Reproducibility: Random seed control and deterministic training for consistent results.\n\n* Validation and Evaluation: Performance monitoring via validation TFRecords and standard regression metrics.\n\nThis work serves as a baseline for large-scale AI models in atmospheric physics, showcasing the efficiency of hybrid frameworks (Keras + JAX) in scientific applications.","metadata":{}},{"cell_type":"markdown","source":"### - Imports & Configuration","metadata":{}},{"cell_type":"code","source":"import os\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\n\nimport gc\nimport numpy as np\nimport pandas as pd\nimport polars as pl\nimport matplotlib.pyplot as plt\nimport tensorflow as tf\nimport jax\nimport keras\nfrom sklearn import metrics\nfrom tqdm.notebook import tqdm\n\nprint(tf.__version__)\nprint(jax.__version__)","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":13.675589,"end_time":"2024-06-17T03:58:21.276994","exception":false,"start_time":"2024-06-17T03:58:07.601405","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:21:51.885715Z","iopub.execute_input":"2024-06-29T08:21:51.886088Z","iopub.status.idle":"2024-06-29T08:22:06.054063Z","shell.execute_reply.started":"2024-06-29T08:21:51.886059Z","shell.execute_reply":"2024-06-29T08:22:06.05316Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def is_interactive():\n    return 'runtime' in get_ipython().config.IPKernelApp.connection_file\n\nprint('Interactive?', is_interactive())","metadata":{"papermill":{"duration":0.015494,"end_time":"2024-06-17T03:58:21.300128","exception":false,"start_time":"2024-06-17T03:58:21.284634","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:06.056021Z","iopub.execute_input":"2024-06-29T08:22:06.056511Z","iopub.status.idle":"2024-06-29T08:22:06.061789Z","shell.execute_reply.started":"2024-06-29T08:22:06.056486Z","shell.execute_reply":"2024-06-29T08:22:06.060804Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 42\nkeras.utils.set_random_seed(SEED)\ntf.random.set_seed(SEED)\ntf.config.experimental.enable_op_determinism()","metadata":{"papermill":{"duration":0.014683,"end_time":"2024-06-17T03:58:21.321988","exception":false,"start_time":"2024-06-17T03:58:21.307305","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:06.063003Z","iopub.execute_input":"2024-06-29T08:22:06.063341Z","iopub.status.idle":"2024-06-29T08:22:06.136935Z","shell.execute_reply.started":"2024-06-29T08:22:06.063309Z","shell.execute_reply":"2024-06-29T08:22:06.135937Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Data Loading  ","metadata":{}},{"cell_type":"code","source":"DATA = \"/kaggle/input/leap-atmospheric-physics-ai-climsim\"\nDATA_TFREC = \"/kaggle/input/leap-train-tfrecords\"","metadata":{"papermill":{"duration":0.013815,"end_time":"2024-06-17T03:58:21.34305","exception":false,"start_time":"2024-06-17T03:58:21.329235","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:06.137897Z","iopub.execute_input":"2024-06-29T08:22:06.138127Z","iopub.status.idle":"2024-06-29T08:22:06.146393Z","shell.execute_reply.started":"2024-06-29T08:22:06.138107Z","shell.execute_reply":"2024-06-29T08:22:06.14563Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample = pl.read_csv(os.path.join(DATA, \"sample_submission.csv\"), n_rows=1)\nTARGETS = sample.select(pl.exclude('sample_id')).columns\nprint(len(TARGETS))","metadata":{"papermill":{"duration":0.211126,"end_time":"2024-06-17T03:58:21.56151","exception":false,"start_time":"2024-06-17T03:58:21.350384","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:06.149051Z","iopub.execute_input":"2024-06-29T08:22:06.149302Z","iopub.status.idle":"2024-06-29T08:22:06.270549Z","shell.execute_reply.started":"2024-06-29T08:22:06.14928Z","shell.execute_reply":"2024-06-29T08:22:06.269591Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def _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']","metadata":{"papermill":{"duration":0.016134,"end_time":"2024-06-17T03:58:21.597012","exception":false,"start_time":"2024-06-17T03:58:21.580878","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:06.271553Z","iopub.execute_input":"2024-06-29T08:22:06.271794Z","iopub.status.idle":"2024-06-29T08:22:06.27674Z","shell.execute_reply.started":"2024-06-29T08:22:06.271773Z","shell.execute_reply":"2024-06-29T08:22:06.275837Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_files = [os.path.join(DATA_TFREC, \"train_%.3d.tfrec\" % i) for i in range(100)]\nvalid_files = [os.path.join(DATA_TFREC, \"train_%.3d.tfrec\" % i) for i in range(100, 101)]","metadata":{"papermill":{"duration":0.015037,"end_time":"2024-06-17T03:58:21.61935","exception":false,"start_time":"2024-06-17T03:58:21.604313","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:06.277729Z","iopub.execute_input":"2024-06-29T08:22:06.278011Z","iopub.status.idle":"2024-06-29T08:22:06.286614Z","shell.execute_reply.started":"2024-06-29T08:22:06.277987Z","shell.execute_reply":"2024-06-29T08:22:06.28586Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. TFRecord Parsing & Dataset Preparation  ","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 4096\n\ntrain_options = tf.data.Options()\ntrain_options.deterministic = True\n\nds_train = (\n    tf.data.Dataset.from_tensor_slices(train_files)\n    .with_options(train_options)\n    .shuffle(100)\n    .interleave(\n        lambda file: tf.data.TFRecordDataset(file).map(_parse_function, 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 * BATCH_SIZE)\n    .batch(BATCH_SIZE)\n    .prefetch(tf.data.AUTOTUNE)\n)\n\nds_valid = (\n    tf.data.TFRecordDataset(valid_files)\n    .map(_parse_function)\n    .batch(BATCH_SIZE)\n    .prefetch(tf.data.AUTOTUNE)\n)","metadata":{"papermill":{"duration":2.17904,"end_time":"2024-06-17T03:58:23.805736","exception":false,"start_time":"2024-06-17T03:58:21.626696","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:06.287622Z","iopub.execute_input":"2024-06-29T08:22:06.287942Z","iopub.status.idle":"2024-06-29T08:22:08.561711Z","shell.execute_reply.started":"2024-06-29T08:22:06.287911Z","shell.execute_reply":"2024-06-29T08:22:08.560895Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Model Architecture ","metadata":{}},{"cell_type":"code","source":"norm_x = keras.layers.Normalization()\nnorm_x.adapt(ds_train.map(lambda x, y: x).take(20 if is_interactive() else 1000))\n\nplt.scatter(\n    norm_x.mean.squeeze(),\n    norm_x.variance.squeeze() ** 0.5,\n    marker=\".\",\n    alpha=0.5\n)\nplt.xscale('log')\nplt.yscale('log')","metadata":{"papermill":{"duration":76.505386,"end_time":"2024-06-17T03:59:40.31865","exception":false,"start_time":"2024-06-17T03:58:23.813264","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:08.562987Z","iopub.execute_input":"2024-06-29T08:22:08.563429Z","iopub.status.idle":"2024-06-29T08:22:13.932487Z","shell.execute_reply.started":"2024-06-29T08:22:08.563394Z","shell.execute_reply":"2024-06-29T08:22:13.931568Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"norm_y = keras.layers.Normalization()\nnorm_y.adapt(ds_train.map(lambda x, y: y).take(20 if is_interactive() else 1000))\n\nmean_y = norm_y.mean\nstdd_y = keras.ops.maximum(1e-10, norm_y.variance ** 0.5)","metadata":{"papermill":{"duration":73.16048,"end_time":"2024-06-17T04:00:53.486936","exception":false,"start_time":"2024-06-17T03:59:40.326456","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:13.933668Z","iopub.execute_input":"2024-06-29T08:22:13.933949Z","iopub.status.idle":"2024-06-29T08:22:17.993731Z","shell.execute_reply.started":"2024-06-29T08:22:13.933923Z","shell.execute_reply":"2024-06-29T08:22:17.99285Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time \n\nmin_y = np.min(np.stack([np.min(yb, 0) for _, yb in ds_train.take(20 if is_interactive() else 1000)], 0), 0, keepdims=True)\nmax_y = np.max(np.stack([np.max(yb, 0) for _, yb in ds_train.take(20 if is_interactive() else 1000)], 0), 0, keepdims=True)","metadata":{"papermill":{"duration":144.709377,"end_time":"2024-06-17T04:03:18.205265","exception":false,"start_time":"2024-06-17T04:00:53.495888","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:17.994972Z","iopub.execute_input":"2024-06-29T08:22:17.995247Z","iopub.status.idle":"2024-06-29T08:22:23.689582Z","shell.execute_reply.started":"2024-06-29T08:22:17.995223Z","shell.execute_reply":"2024-06-29T08:22:23.688522Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### - Model definition & Training","metadata":{"papermill":{"duration":0.007927,"end_time":"2024-06-17T04:03:18.222885","exception":false,"start_time":"2024-06-17T04:03:18.214958","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":0.018682,"end_time":"2024-06-17T04:03:18.249621","exception":false,"start_time":"2024-06-17T04:03:18.230939","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:23.690677Z","iopub.execute_input":"2024-06-29T08:22:23.690962Z","iopub.status.idle":"2024-06-29T08:22:23.698324Z","shell.execute_reply.started":"2024-06-29T08:22:23.690937Z","shell.execute_reply":"2024-06-29T08:22:23.697425Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"steps_per_epoch = int(np.ceil(len(train_files) * 100_000 / BATCH_SIZE))\nprint(steps_per_epoch)","metadata":{"execution":{"iopub.status.busy":"2024-06-29T08:26:47.105859Z","iopub.execute_input":"2024-06-29T08:26:47.106593Z","iopub.status.idle":"2024-06-29T08:26:47.111384Z","shell.execute_reply.started":"2024-06-29T08:26:47.10656Z","shell.execute_reply":"2024-06-29T08:26:47.110483Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"epochs = 15 \nlearning_rate = 1e-3\n\nepochs_warmup = 2\nepochs_ending = 3\nsteps_per_epoch = int(np.ceil(len(train_files) * 100_000 / BATCH_SIZE))\n\nlr_scheduler = keras.optimizers.schedules.CosineDecay(\n    1e-4, \n    (epochs - epochs_warmup - epochs_ending) * steps_per_epoch, \n    warmup_target=learning_rate,\n    warmup_steps=steps_per_epoch * epochs_warmup,\n    alpha=0.1\n)\n\nplt.plot([lr_scheduler(it) for it in range(0, epochs * steps_per_epoch, steps_per_epoch)]);","metadata":{"papermill":{"duration":4.770512,"end_time":"2024-06-17T04:03:23.028441","exception":false,"start_time":"2024-06-17T04:03:18.257929","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:27:19.494495Z","iopub.execute_input":"2024-06-29T08:27:19.495289Z","iopub.status.idle":"2024-06-29T08:27:20.188906Z","shell.execute_reply.started":"2024-06-29T08:27:19.49526Z","shell.execute_reply":"2024-06-29T08:27:20.187944Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"keras.utils.clear_session()\n\n\ndef 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\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\n\nX_input = x = keras.layers.Input(ds_train.element_spec[0].shape[1:])\nx = keras.layers.Normalization(mean=norm_x.mean, variance=norm_x.variance)(x)\nx = x_to_seq(x)\n\n\ne = e0 = keras.layers.Conv1D(64, 1, padding='same')(x)\ne = build_cnn()(e)\n# add global average to allow some comunication between all levels even in a small CNN\ne = e0 + e + keras.layers.GlobalAveragePooling1D(keepdims=True)(e)\ne = keras.layers.BatchNormalization()(e)\ne = e + build_cnn()(e)\n\n\np_all = keras.layers.Conv1D(14, 1, padding='same')(e)\n\np_seq = p_all[:, :, :6]\np_seq = keras.ops.transpose(p_seq, (0, 2, 1))\np_seq = keras.layers.Flatten()(p_seq)\nassert p_seq.shape[-1] == 360\n\np_flat = p_all[:, :, 6:6 + 8]\np_flat = keras.ops.mean(p_flat, axis=1)\nassert p_flat.shape[-1] == 8\n\nP = keras.ops.concatenate([p_seq, p_flat], axis=1)\n\n# build & compile\nmodel = keras.Model(X_input, P)\nmodel.compile(\n    loss='mse', \n    optimizer=keras.optimizers.Adam(lr_scheduler),\n    metrics=[ClippedR2Score()]\n)\nmodel.build(tuple(ds_train.element_spec[0].shape))\nmodel.summary()","metadata":{"papermill":{"duration":1.108948,"end_time":"2024-06-17T04:03:24.145981","exception":false,"start_time":"2024-06-17T04:03:23.037033","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:27:27.127191Z","iopub.execute_input":"2024-06-29T08:27:27.128035Z","iopub.status.idle":"2024-06-29T08:27:28.270117Z","shell.execute_reply.started":"2024-06-29T08:27:27.128003Z","shell.execute_reply":"2024-06-29T08:27:28.269276Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Model Training  ","metadata":{}},{"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))\n\nhistory = model.fit(\n    ds_train_target_normalized,\n    validation_data=ds_valid_target_normalized,\n    epochs=epochs,\n    verbose=1 if is_interactive() else 2,\n    callbacks=[\n        keras.callbacks.ModelCheckpoint(filepath='model.keras')\n    ]\n)","metadata":{"papermill":{"duration":28930.893576,"end_time":"2024-06-17T12:05:35.049837","exception":false,"start_time":"2024-06-17T04:03:24.156261","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:27:33.115115Z","iopub.execute_input":"2024-06-29T08:27:33.115586Z","iopub.status.idle":"2024-06-29T09:07:51.43082Z","shell.execute_reply.started":"2024-06-29T08:27:33.115554Z","shell.execute_reply":"2024-06-29T09:07:51.426761Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6. Evaluation","metadata":{}},{"cell_type":"code","source":"plt.plot(history.history['loss'], color='tab:blue')\nplt.plot(history.history['val_loss'], color='tab:red')\nplt.yscale('log');\n\ny_valid = np.concatenate([yb for _, yb in ds_valid])\np_valid = model.predict(ds_valid, batch_size=BATCH_SIZE) * stdd_y + mean_y\n\nscores_valid = np.array([metrics.r2_score(y_valid[:, i], p_valid[:, i]) for i in range(len(TARGETS))])\nplt.plot(scores_valid.clip(-1, 1))\n\nmask = scores_valid <= 1e-3\nf\"Number of under-performing targets: {sum(mask)}\"\n\nf\"Clipped score: {scores_valid.clip(0, 1).mean()}\"\n\ndel y_valid, p_valid\ngc.collect();","metadata":{"papermill":{"duration":0.387267,"end_time":"2024-06-17T12:05:35.462488","exception":false,"start_time":"2024-06-17T12:05:35.075221","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T09:12:58.630712Z","iopub.execute_input":"2024-06-29T09:12:58.631742Z","iopub.status.idle":"2024-06-29T09:12:59.056956Z","shell.execute_reply.started":"2024-06-29T09:12:58.631703Z","shell.execute_reply":"2024-06-29T09:12:59.056075Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 7. Submission","metadata":{"papermill":{"duration":0.02795,"end_time":"2024-06-17T12:05:43.020167","exception":false,"start_time":"2024-06-17T12:05:42.992217","status":"completed"},"tags":[]}},{"cell_type":"code","source":"sample = pl.read_csv(\"/kaggle/input/leap-atmospheric-physics-ai-climsim/sample_submission.csv\")","metadata":{"papermill":{"duration":16.162126,"end_time":"2024-06-17T12:05:59.209904","exception":false,"start_time":"2024-06-17T12:05:43.047778","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T09:13:26.075907Z","iopub.execute_input":"2024-06-29T09:13:26.076267Z","iopub.status.idle":"2024-06-29T09:13:29.865709Z","shell.execute_reply.started":"2024-06-29T09:13:26.076237Z","shell.execute_reply":"2024-06-29T09:13:29.864865Z"},"trusted":true},"outputs":[],"execution_count":null},{"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)","metadata":{"papermill":{"duration":30.76129,"end_time":"2024-06-17T12:06:30","exception":false,"start_time":"2024-06-17T12:05:59.23871","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T09:13:32.765198Z","iopub.execute_input":"2024-06-29T09:13:32.766064Z","iopub.status.idle":"2024-06-29T09:14:01.873314Z","shell.execute_reply.started":"2024-06-29T09:13:32.766029Z","shell.execute_reply":"2024-06-29T09:14:01.870334Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"p_test = model.predict(df_test.to_numpy(), batch_size=4 * BATCH_SIZE) * stdd_y + mean_y\np_test = np.array(p_test)\np_test[:, mask] = mean_y[:, mask]","metadata":{"papermill":{"duration":107.247373,"end_time":"2024-06-17T12:08:17.282886","exception":false,"start_time":"2024-06-17T12:06:30.035513","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T09:14:12.046995Z","iopub.execute_input":"2024-06-29T09:14:12.047389Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# correction of ptend_q0002 targets (from 12 to 29)\ndf_p_test = pd.DataFrame(p_test, columns=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":{"papermill":{"duration":42.181472,"end_time":"2024-06-17T12:08:59.495014","exception":false,"start_time":"2024-06-17T12:08:17.313542","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:24.489244Z","iopub.status.idle":"2024-06-29T08:22:24.48972Z","shell.execute_reply.started":"2024-06-29T08:22:24.489492Z","shell.execute_reply":"2024-06-29T08:22:24.489511Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = sample.to_pandas()\nsubmission[TARGETS] = submission[TARGETS] * p_test\npl.from_pandas(submission[[\"sample_id\"] + TARGETS]).write_csv(\"submission.csv\")","metadata":{"papermill":{"duration":154.388028,"end_time":"2024-06-17T12:11:33.91407","exception":false,"start_time":"2024-06-17T12:08:59.526042","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-06-29T08:22:24.491047Z","iopub.status.idle":"2024-06-29T08:22:24.491509Z","shell.execute_reply.started":"2024-06-29T08:22:24.491255Z","shell.execute_reply":"2024-06-29T08:22:24.491273Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 8. Conclusion\n\nThe experiment demonstrates that deep neural networks trained with Keras (JAX backend) can effectively learn complex atmospheric dynamics from the LEAP Climsim dataset.\n\nKey insights:\n\n* The TFRecord pipeline significantly improves data throughput and GPU utilization during training.\n\n* Deterministic setup (controlled seeds, fixed batch operations) ensures reproducible results — crucial for scientific benchmarks.\n\n* The multi-output Keras model achieves stable convergence across 368 target variables, highlighting the representational power of neural networks for high-dimensional climate prediction.\n\nFuture directions:\n\n* Experiment with transformer-based architectures for spatiotemporal generalization.\n\n* Apply physics-informed loss functions to enforce conservation laws.\n\n* Explore transfer learning across different atmospheric regimes.\n\nThis notebook establishes a strong and efficient foundation for AI-driven climate modeling — combining physical insight with modern machine learning practices.","metadata":{}}]}