{"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":8015876,"sourceType":"competition"},{"sourceId":8409068,"sourceType":"datasetVersion","datasetId":5004471}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\n\nimport tensorflow as tf\nimport jax\nimport keras\n\nfrom sklearn import metrics\n\nfrom tqdm.notebook import tqdm\n\nprint(tf.__version__)\nprint(jax.__version__)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-29T08:17:27.121495Z","iopub.execute_input":"2024-05-29T08:17:27.121868Z","iopub.status.idle":"2024-05-29T08:17:44.593706Z","shell.execute_reply.started":"2024-05-29T08:17:27.121838Z","shell.execute_reply":"2024-05-29T08:17:44.592079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def is_interactive():\n    return 'runtime' in get_ipython().config.IPKernelApp.connection_file\n\nprint('Interactive?', is_interactive())","metadata":{"execution":{"iopub.status.busy":"2024-05-29T08:17:44.606348Z","iopub.execute_input":"2024-05-29T08:17:44.60702Z","iopub.status.idle":"2024-05-29T08:17:44.613986Z","shell.execute_reply.started":"2024-05-29T08:17:44.606983Z","shell.execute_reply":"2024-05-29T08:17:44.612557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 42\nkeras.utils.set_random_seed(SEED)\ntf.random.set_seed(SEED)\ntf.config.experimental.enable_op_determinism()","metadata":{"execution":{"iopub.status.busy":"2024-05-29T08:17:44.615314Z","iopub.execute_input":"2024-05-29T08:17:44.615654Z","iopub.status.idle":"2024-05-29T08:17:44.642532Z","shell.execute_reply.started":"2024-05-29T08:17:44.615625Z","shell.execute_reply":"2024-05-29T08:17:44.63877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA = \"/kaggle/input/leap-atmospheric-physics-ai-climsim\"\nDATA_TFREC = \"/kaggle/input/leap-train-tfrecords\"","metadata":{"execution":{"iopub.status.busy":"2024-05-29T08:17:44.644672Z","iopub.execute_input":"2024-05-29T08:17:44.645841Z","iopub.status.idle":"2024-05-29T08:17:44.657667Z","shell.execute_reply.started":"2024-05-29T08:17:44.645794Z","shell.execute_reply":"2024-05-29T08:17:44.656193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-05-29T08:17:44.659146Z","iopub.execute_input":"2024-05-29T08:17:44.65957Z","iopub.status.idle":"2024-05-29T08:17:44.880207Z","shell.execute_reply.started":"2024-05-29T08:17:44.659531Z","shell.execute_reply":"2024-05-29T08:17:44.878775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-05-29T08:17:44.881812Z","iopub.execute_input":"2024-05-29T08:17:44.882266Z","iopub.status.idle":"2024-05-29T08:17:44.890949Z","shell.execute_reply.started":"2024-05-29T08:17:44.882204Z","shell.execute_reply":"2024-05-29T08:17:44.889439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-05-29T08:17:44.892992Z","iopub.execute_input":"2024-05-29T08:17:44.893474Z","iopub.status.idle":"2024-05-29T08:17:44.909201Z","shell.execute_reply.started":"2024-05-29T08:17:44.893434Z","shell.execute_reply":"2024-05-29T08:17:44.907815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 2048\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":{"execution":{"iopub.status.busy":"2024-05-29T08:17:44.913696Z","iopub.execute_input":"2024-05-29T08:17:44.914174Z","iopub.status.idle":"2024-05-29T08:17:45.172709Z","shell.execute_reply.started":"2024-05-29T08:17:44.914127Z","shell.execute_reply":"2024-05-29T08:17:45.171375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-05-29T08:17:45.17476Z","iopub.execute_input":"2024-05-29T08:17:45.17585Z","iopub.status.idle":"2024-05-29T08:17:52.134645Z","shell.execute_reply.started":"2024-05-29T08:17:45.175797Z","shell.execute_reply":"2024-05-29T08:17:52.133293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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)\n\nplt.scatter(\n    mean_y.squeeze(),\n    stdd_y.squeeze(),\n    marker=\".\",\n    alpha=0.5\n)\nplt.xscale('log')\nplt.yscale('log')\n","metadata":{"execution":{"iopub.status.busy":"2024-05-29T08:17:52.136691Z","iopub.execute_input":"2024-05-29T08:17:52.13715Z","iopub.status.idle":"2024-05-29T08:17:57.276043Z","shell.execute_reply.started":"2024-05-29T08:17:52.137105Z","shell.execute_reply":"2024-05-29T08:17:57.274956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"min_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":{"execution":{"iopub.status.busy":"2024-05-29T08:17:57.277338Z","iopub.execute_input":"2024-05-29T08:17:57.277642Z","iopub.status.idle":"2024-05-29T08:18:04.925632Z","shell.execute_reply.started":"2024-05-29T08:17:57.277615Z","shell.execute_reply":"2024-05-29T08:18:04.924343Z"},"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_count":null,"outputs":[]},{"cell_type":"code","source":"epochs = 55 # 25  # 15  # 12\nlearning_rate = 1e-3\n\nepochs_warmup = 1\nepochs_ending = 2\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":{"execution":{"iopub.status.busy":"2024-05-29T08:18:04.927967Z","iopub.execute_input":"2024-05-29T08:18:04.928417Z","iopub.status.idle":"2024-05-29T08:18:05.634824Z","shell.execute_reply.started":"2024-05-29T08:18:04.928367Z","shell.execute_reply":"2024-05-29T08:18:05.633676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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        # Première couche Conv1D\n        keras.layers.Conv1D(256, 3, padding='same', activation=activation),\n        keras.layers.BatchNormalization(),\n        # Deuxième couche Conv1D\n        keras.layers.Conv1D(128, 3, padding='same', activation=activation),\n        keras.layers.BatchNormalization(),\n        # Troisième couche Conv1D\n        keras.layers.Conv1D(64, 3, padding='same', activation=activation),\n        keras.layers.BatchNormalization(),\n        # Couche Dropout pour la régularisation\n        keras.layers.Dropout(0.3),\n        # Couche LSTM pour capturer les dépendances temporelles\n        keras.layers.LSTM(64, return_sequences=True),\n        keras.layers.BatchNormalization(),\n        # Couche GRU pour capturer les dépendances temporelles\n        keras.layers.GRU(64, return_sequences=True),\n        keras.layers.BatchNormalization(),\n        # Couche Dense finale pour ajuster la sortie à la taille 64\n        keras.layers.Dense(64, activation=activation),\n    ])\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":{"execution":{"iopub.status.busy":"2024-05-29T08:29:07.000656Z","iopub.execute_input":"2024-05-29T08:29:07.001275Z","iopub.status.idle":"2024-05-29T08:29:07.771468Z","shell.execute_reply.started":"2024-05-29T08:29:07.001207Z","shell.execute_reply":"2024-05-29T08:29:07.77025Z"},"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))\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":{"execution":{"iopub.status.busy":"2024-05-29T08:29:09.921305Z","iopub.execute_input":"2024-05-29T08:29:09.921732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(history.history['loss'], color='tab:blue')\nplt.plot(history.history['val_loss'], color='tab:red')\nplt.yscale('log');","metadata":{"execution":{"iopub.status.busy":"2024-05-19T22:02:36.921199Z","iopub.execute_input":"2024-05-19T22:02:36.92188Z","iopub.status.idle":"2024-05-19T22:02:36.959564Z","shell.execute_reply.started":"2024-05-19T22:02:36.921843Z","shell.execute_reply":"2024-05-19T22:02:36.958285Z"},"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=BATCH_SIZE) * stdd_y + mean_y","metadata":{"execution":{"iopub.status.busy":"2024-05-19T22:02:57.618947Z","iopub.execute_input":"2024-05-19T22:02:57.619914Z","iopub.status.idle":"2024-05-19T22:03:02.431339Z","shell.execute_reply.started":"2024-05-19T22:02:57.619877Z","shell.execute_reply":"2024-05-19T22:03:02.430321Z"},"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(TARGETS))])\nplt.plot(scores_valid.clip(-1, 1))","metadata":{"execution":{"iopub.status.busy":"2024-05-19T22:03:04.551666Z","iopub.execute_input":"2024-05-19T22:03:04.552027Z","iopub.status.idle":"2024-05-19T22:03:06.007572Z","shell.execute_reply.started":"2024-05-19T22:03:04.552002Z","shell.execute_reply":"2024-05-19T22:03:06.00668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = scores_valid <= 1e-3\nf\"Number of under-performing targets: {sum(mask)}\"","metadata":{"execution":{"iopub.status.busy":"2024-05-19T22:03:10.92824Z","iopub.execute_input":"2024-05-19T22:03:10.928992Z","iopub.status.idle":"2024-05-19T22:03:10.935136Z","shell.execute_reply.started":"2024-05-19T22:03:10.928958Z","shell.execute_reply":"2024-05-19T22:03:10.934227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f\"Clipped score: {scores_valid.clip(0, 1).mean()}\"","metadata":{"execution":{"iopub.status.busy":"2024-05-19T22:03:14.882563Z","iopub.execute_input":"2024-05-19T22:03:14.883709Z","iopub.status.idle":"2024-05-19T22:03:14.889832Z","shell.execute_reply.started":"2024-05-19T22:03:14.883664Z","shell.execute_reply":"2024-05-19T22:03:14.888863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del y_valid, p_valid\ngc.collect();","metadata":{"execution":{"iopub.status.busy":"2024-05-19T22:03:17.095698Z","iopub.execute_input":"2024-05-19T22:03:17.096537Z","iopub.status.idle":"2024-05-19T22:03:17.476466Z","shell.execute_reply.started":"2024-05-19T22:03:17.096504Z","shell.execute_reply":"2024-05-19T22:03:17.475659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"sample = pl.read_csv(\"/kaggle/input/leap-atmospheric-physics-ai-climsim/sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-05-19T22:03:23.312602Z","iopub.execute_input":"2024-05-19T22:03:23.313512Z","iopub.status.idle":"2024-05-19T22:03:37.916699Z","shell.execute_reply.started":"2024-05-19T22:03:23.313471Z","shell.execute_reply":"2024-05-19T22:03:37.915682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-05-19T22:03:37.918373Z","iopub.execute_input":"2024-05-19T22:03:37.918676Z","iopub.status.idle":"2024-05-19T22:04:04.359645Z","shell.execute_reply.started":"2024-05-19T22:03:37.918644Z","shell.execute_reply":"2024-05-19T22:04:04.358605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-05-19T22:04:04.360892Z","iopub.execute_input":"2024-05-19T22:04:04.361157Z","iopub.status.idle":"2024-05-19T22:04:12.652983Z","shell.execute_reply.started":"2024-05-19T22:04:04.361135Z","shell.execute_reply":"2024-05-19T22:04:12.652138Z"},"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=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":{"execution":{"iopub.status.busy":"2024-05-19T22:04:12.654833Z","iopub.execute_input":"2024-05-19T22:04:12.655147Z","iopub.status.idle":"2024-05-19T22:04:16.104079Z","shell.execute_reply.started":"2024-05-19T22:04:12.655123Z","shell.execute_reply":"2024-05-19T22:04:16.101269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-05-19T22:04:16.108747Z","iopub.execute_input":"2024-05-19T22:04:16.109588Z","iopub.status.idle":"2024-05-19T22:04:43.132139Z","shell.execute_reply.started":"2024-05-19T22:04:16.109511Z","shell.execute_reply":"2024-05-19T22:04:43.131288Z"},"trusted":true},"execution_count":null,"outputs":[]}]}