{"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":30698,"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-19T21:56:59.935406Z","iopub.execute_input":"2024-05-19T21:56:59.935876Z","iopub.status.idle":"2024-05-19T21:57:12.778252Z","shell.execute_reply.started":"2024-05-19T21:56:59.935834Z","shell.execute_reply":"2024-05-19T21:57:12.7773Z"},"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-19T21:57:12.780266Z","iopub.execute_input":"2024-05-19T21:57:12.781017Z","iopub.status.idle":"2024-05-19T21:57:12.786017Z","shell.execute_reply.started":"2024-05-19T21:57:12.780981Z","shell.execute_reply":"2024-05-19T21:57:12.784974Z"},"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-19T21:57:12.787426Z","iopub.execute_input":"2024-05-19T21:57:12.787771Z","iopub.status.idle":"2024-05-19T21:57:12.812279Z","shell.execute_reply.started":"2024-05-19T21:57:12.78774Z","shell.execute_reply":"2024-05-19T21:57:12.811261Z"},"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-19T21:57:12.814398Z","iopub.execute_input":"2024-05-19T21:57:12.814689Z","iopub.status.idle":"2024-05-19T21:57:12.822873Z","shell.execute_reply.started":"2024-05-19T21:57:12.814665Z","shell.execute_reply":"2024-05-19T21:57:12.822108Z"},"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-19T21:57:12.823892Z","iopub.execute_input":"2024-05-19T21:57:12.824193Z","iopub.status.idle":"2024-05-19T21:57:13.013907Z","shell.execute_reply.started":"2024-05-19T21:57:12.82417Z","shell.execute_reply":"2024-05-19T21:57:13.012958Z"},"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-19T21:57:13.015117Z","iopub.execute_input":"2024-05-19T21:57:13.015777Z","iopub.status.idle":"2024-05-19T21:57:13.020952Z","shell.execute_reply.started":"2024-05-19T21:57:13.015743Z","shell.execute_reply":"2024-05-19T21:57:13.020054Z"},"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-19T21:57:13.022056Z","iopub.execute_input":"2024-05-19T21:57:13.02231Z","iopub.status.idle":"2024-05-19T21:57:13.030757Z","shell.execute_reply.started":"2024-05-19T21:57:13.022289Z","shell.execute_reply":"2024-05-19T21:57:13.029943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-05-19T21:57:13.031848Z","iopub.execute_input":"2024-05-19T21:57:13.032075Z","iopub.status.idle":"2024-05-19T21:57:15.128056Z","shell.execute_reply.started":"2024-05-19T21:57:13.032055Z","shell.execute_reply":"2024-05-19T21:57:15.127219Z"},"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-19T21:57:15.129171Z","iopub.execute_input":"2024-05-19T21:57:15.129417Z","iopub.status.idle":"2024-05-19T21:57:21.13128Z","shell.execute_reply.started":"2024-05-19T21:57:15.129396Z","shell.execute_reply":"2024-05-19T21:57:21.130367Z"},"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-19T21:57:21.135044Z","iopub.execute_input":"2024-05-19T21:57:21.135341Z","iopub.status.idle":"2024-05-19T21:57:25.35317Z","shell.execute_reply.started":"2024-05-19T21:57:21.135315Z","shell.execute_reply":"2024-05-19T21:57:25.352274Z"},"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-19T21:57:25.354336Z","iopub.execute_input":"2024-05-19T21:57:25.35462Z","iopub.status.idle":"2024-05-19T21:57:31.171634Z","shell.execute_reply.started":"2024-05-19T21:57:25.354582Z","shell.execute_reply":"2024-05-19T21:57:31.170856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model definition & Training","metadata":{}},{"cell_type":"code","source":"epochs = 10\nlearning_rate = 1e-3\n\nepochs_warmup = 2\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-19T21:57:31.17274Z","iopub.execute_input":"2024-05-19T21:57:31.173023Z","iopub.status.idle":"2024-05-19T21:57:32.538772Z","shell.execute_reply.started":"2024-05-19T21:57:31.172998Z","shell.execute_reply":"2024-05-19T21:57:32.537593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = keras.Sequential([\n    keras.layers.Normalization(mean=norm_x.mean, variance=norm_x.variance),\n    \n    keras.layers.Dense(1024, activation='relu'),\n    keras.layers.Dense(512, activation='relu'),\n    \n    keras.layers.Dense(len(TARGETS))\n])\nmodel.compile(loss='mse', optimizer=keras.optimizers.Adam(lr_scheduler))\nmodel.build(tuple(ds_train.element_spec[0].shape))\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2024-05-19T22:02:53.847922Z","iopub.execute_input":"2024-05-19T22:02:53.848613Z","iopub.status.idle":"2024-05-19T22:02:53.882409Z","shell.execute_reply.started":"2024-05-19T22:02:53.848578Z","shell.execute_reply":"2024-05-19T22:02:53.881521Z"},"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-19T21:57:33.250185Z","iopub.execute_input":"2024-05-19T21:57:33.250459Z","iopub.status.idle":"2024-05-19T22:02:34.26318Z","shell.execute_reply.started":"2024-05-19T21:57:33.250434Z","shell.execute_reply":"2024-05-19T22:02:34.261765Z"},"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":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}