{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":56537,"databundleVersionId":8015876,"sourceType":"competition"},{"sourceId":8409068,"sourceType":"datasetVersion","datasetId":5004471},{"sourceId":8652896,"sourceType":"datasetVersion","datasetId":5183232}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!pip uninstall -y tensorflow\n#!pip install --upgrade tensorflow","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-06-14T14:20:05.183864Z","iopub.execute_input":"2024-06-14T14:20:05.184212Z","iopub.status.idle":"2024-06-14T14:20:05.192726Z","shell.execute_reply.started":"2024-06-14T14:20:05.184182Z","shell.execute_reply":"2024-06-14T14:20:05.19195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.environ[\"KERAS_BACKEND\"] = \"tensorflow\"\n\nimport gc\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\n#import jax\nimport tensorflow as tf\nfrom tensorflow import keras\n\nfrom sklearn import metrics\n\nfrom tqdm.notebook import tqdm\n\nprint(tf.__version__)\n#print(jax.__version__)","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:20:05.194146Z","iopub.execute_input":"2024-06-14T14:20:05.194411Z","iopub.status.idle":"2024-06-14T14:20:23.38736Z","shell.execute_reply.started":"2024-06-14T14:20:05.194367Z","shell.execute_reply":"2024-06-14T14:20:23.386661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install polars","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:20:23.388223Z","iopub.execute_input":"2024-06-14T14:20:23.388659Z","iopub.status.idle":"2024-06-14T14:20:30.093461Z","shell.execute_reply.started":"2024-06-14T14:20:23.38863Z","shell.execute_reply":"2024-06-14T14:20:30.092304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import polars as pl","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:20:30.09486Z","iopub.execute_input":"2024-06-14T14:20:30.095175Z","iopub.status.idle":"2024-06-14T14:20:30.603499Z","shell.execute_reply.started":"2024-06-14T14:20:30.095142Z","shell.execute_reply":"2024-06-14T14:20:30.602692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_device():\n    \"Detect and intializes GPU/TPU automatically\"\n    # Check TPU category\n    tpu = 'local'\n    try:\n        # Connect to TPU\n        print(\"Connecting to TPU...\")\n        tpu = tf.distribute.cluster_resolver.TPUClusterResolver() \n        tf.tpu.experimental.initialize_tpu_system(tpu)\n        #tpu = tf.distribute.cluster_resolver.TPUClusterResolver.connect(tpu=tpu) \n        # Set TPU strategy\n        print(\"Setting TPU ste\")\n        strategy = tf.distribute.TPUStrategy(tpu)\n        print(f'> Running on TPU ', tpu.master(), end=' | ')\n        print('Num of TPUs: ', strategy.num_replicas_in_sync)\n        device='TPU'\n    except:\n        print(\"TPU not available!\")\n        # If TPU is not available, detect GPUs\n        gpus = tf.config.list_logical_devices('GPU')\n        ngpu = len(gpus)\n         # Check number of GPUs\n        if ngpu:\n            # Set GPU strategy\n            strategy = tf.distribute.MirroredStrategy(gpus) # single-GPU or multi-GPU\n            # Print GPU details\n            print(\"> Running on GPU\", end=' | ')\n            print(\"Num of GPUs: \", ngpu)\n            device='GPU'\n        else:\n            # If no GPUs are available, use CPU\n            print(\"> Running on CPU\")\n            strategy = tf.distribute.get_strategy()\n            device='CPU'\n    return strategy, device, tpu","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:20:30.605526Z","iopub.execute_input":"2024-06-14T14:20:30.605775Z","iopub.status.idle":"2024-06-14T14:20:30.612054Z","shell.execute_reply.started":"2024-06-14T14:20:30.60575Z","shell.execute_reply":"2024-06-14T14:20:30.611427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tpu_strategy, device, tpu = get_device()\nreplicas = tpu_strategy.num_replicas_in_sync","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:20:30.612985Z","iopub.execute_input":"2024-06-14T14:20:30.61324Z","iopub.status.idle":"2024-06-14T14:20:39.674296Z","shell.execute_reply.started":"2024-06-14T14:20:30.613212Z","shell.execute_reply":"2024-06-14T14:20:39.673428Z"},"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-06-14T14:20:39.675407Z","iopub.execute_input":"2024-06-14T14:20:39.675688Z","iopub.status.idle":"2024-06-14T14:20:39.67974Z","shell.execute_reply.started":"2024-06-14T14:20:39.675658Z","shell.execute_reply":"2024-06-14T14:20:39.679054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 42\ntf.keras.utils.set_random_seed(SEED)\ntf.random.set_seed(SEED)\ntf.config.experimental.enable_op_determinism()","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:20:39.680503Z","iopub.execute_input":"2024-06-14T14:20:39.680716Z","iopub.status.idle":"2024-06-14T14:20:39.723512Z","shell.execute_reply.started":"2024-06-14T14:20:39.680693Z","shell.execute_reply":"2024-06-14T14:20:39.722787Z"},"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\"\nSTATS = \"/kaggle/input/leap-statistics\"","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:20:39.724431Z","iopub.execute_input":"2024-06-14T14:20:39.724683Z","iopub.status.idle":"2024-06-14T14:20:39.727803Z","shell.execute_reply.started":"2024-06-14T14:20:39.724656Z","shell.execute_reply":"2024-06-14T14:20:39.727173Z"},"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-06-14T14:20:39.728604Z","iopub.execute_input":"2024-06-14T14:20:39.728845Z","iopub.status.idle":"2024-06-14T14:20:39.867423Z","shell.execute_reply.started":"2024-06-14T14:20:39.728821Z","shell.execute_reply":"2024-06-14T14:20:39.866672Z"},"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-06-14T14:20:39.868185Z","iopub.execute_input":"2024-06-14T14:20:39.868426Z","iopub.status.idle":"2024-06-14T14:20:39.872473Z","shell.execute_reply.started":"2024-06-14T14:20:39.868401Z","shell.execute_reply":"2024-06-14T14:20:39.871811Z"},"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-06-14T14:20:39.873321Z","iopub.execute_input":"2024-06-14T14:20:39.873658Z","iopub.status.idle":"2024-06-14T14:20:39.882281Z","shell.execute_reply.started":"2024-06-14T14:20:39.873629Z","shell.execute_reply":"2024-06-14T14:20:39.881649Z"},"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-06-14T14:20:39.883102Z","iopub.execute_input":"2024-06-14T14:20:39.883376Z","iopub.status.idle":"2024-06-14T14:20:40.011692Z","shell.execute_reply.started":"2024-06-14T14:20:39.883349Z","shell.execute_reply":"2024-06-14T14:20:40.011015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#norm_x = keras.layers.Normalization()\n#norm_x.adapt(ds_train.map(lambda x, y: x).take(20 if is_interactive() else 1000))\n\nmean_x = np.load(f\"{STATS}/mean_x.npy\")\nvar_x = np.load(f\"{STATS}/var_x.npy\")\n\nplt.scatter(\n    np.squeeze(mean_x),\n    np.squeeze(var_x) ** 0.5,\n    marker=\".\",\n    alpha=0.5\n)\nplt.xscale('log')\nplt.yscale('log')","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:20:40.014707Z","iopub.execute_input":"2024-06-14T14:20:40.014978Z","iopub.status.idle":"2024-06-14T14:20:40.453131Z","shell.execute_reply.started":"2024-06-14T14:20:40.01495Z","shell.execute_reply":"2024-06-14T14:20:40.452424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#norm_y = keras.layers.Normalization()\n#norm_y.adapt(ds_train.map(lambda x, y: y).take(20 if is_interactive() else 1000))\n\nmean_y = np.load(f\"{STATS}/mean_y.npy\")\nvar_y = np.load(f\"{STATS}/var_y.npy\")\n\nstdd_y = np.maximum(1e-10, var_y ** 0.5)\n\nplt.scatter(\n    np.squeeze(mean_y),\n    np.squeeze(stdd_y),\n    marker=\".\",\n    alpha=0.5\n)\nplt.xscale('log')\nplt.yscale('log')\n","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:20:40.454142Z","iopub.execute_input":"2024-06-14T14:20:40.45445Z","iopub.status.idle":"2024-06-14T14:20:40.709586Z","shell.execute_reply.started":"2024-06-14T14:20:40.454418Z","shell.execute_reply":"2024-06-14T14:20:40.708884Z"},"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-06-14T14:20:40.710552Z","iopub.execute_input":"2024-06-14T14:20:40.710802Z","iopub.status.idle":"2024-06-14T14:20:44.920283Z","shell.execute_reply.started":"2024-06-14T14:20:40.710776Z","shell.execute_reply":"2024-06-14T14:20:44.918968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model definition & Training","metadata":{}},{"cell_type":"code","source":"\n@keras.utils.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 tf.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()\n","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:20:44.921749Z","iopub.execute_input":"2024-06-14T14:20:44.922106Z","iopub.status.idle":"2024-06-14T14:20:44.929717Z","shell.execute_reply.started":"2024-06-14T14:20:44.92207Z","shell.execute_reply":"2024-06-14T14:20:44.928779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"epochs = 100  # 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\nwith tpu_strategy.scope():\n    lr_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\n    plt.plot([lr_scheduler(it) for it in range(0, epochs * steps_per_epoch, steps_per_epoch)]);","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:22:25.867092Z","iopub.execute_input":"2024-06-14T14:22:25.867521Z","iopub.status.idle":"2024-06-14T14:22:26.507523Z","shell.execute_reply.started":"2024-06-14T14:22:25.867485Z","shell.execute_reply":"2024-06-14T14:22:26.506697Z"},"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        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\nwith tpu_strategy.scope():\n    X_input = x = keras.layers.Input(ds_train.element_spec[0].shape[1:])\n    x = keras.layers.Normalization(mean=mean_x, variance=var_x)(x)\n    x = x_to_seq(x)\n\n\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\n\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\n    model = keras.Model(X_input, P)\n    \n    model.compile(\n        loss='mse', \n        optimizer=keras.optimizers.Adam(lr_scheduler),\n        metrics=[ClippedR2Score()]\n    )\n    model.build(tuple(ds_train.element_spec[0].shape))\n    model.summary()","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:22:26.509015Z","iopub.execute_input":"2024-06-14T14:22:26.509297Z","iopub.status.idle":"2024-06-14T14:22:27.130187Z","shell.execute_reply.started":"2024-06-14T14:22:26.509267Z","shell.execute_reply":"2024-06-14T14:22:27.129415Z"},"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\nwith tpu_strategy.scope():\n    history = 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-06-14T14:22:27.131093Z","iopub.execute_input":"2024-06-14T14:22:27.131341Z"},"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":{"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":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = scores_valid <= 1e-3\nf\"Number of under-performing targets: {sum(mask)}\"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f\"Clipped score: {scores_valid.clip(0, 1).mean()}\"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del y_valid, p_valid\ngc.collect();","metadata":{"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":{"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":{"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":{"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":{"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":{"trusted":true},"execution_count":null,"outputs":[]}]}