{"metadata":{"accelerator":"TPU","colab":{"gpuType":"V28","machine_shape":"hm","name":"","version":""},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":56537,"databundleVersionId":8877088,"sourceType":"competition"},{"sourceId":8523424,"sourceType":"datasetVersion","datasetId":5085016},{"sourceId":8638333,"sourceType":"datasetVersion","datasetId":5099933},{"sourceId":8874581,"sourceType":"datasetVersion","datasetId":4945873},{"sourceId":8998766,"sourceType":"datasetVersion","datasetId":5416189},{"sourceId":189085504,"sourceType":"kernelVersion"},{"sourceId":189085508,"sourceType":"kernelVersion"}],"dockerImageVersionId":30683,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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"},"papermill":{"default_parameters":{},"duration":10445.047101,"end_time":"2024-06-06T17:03:34.206193","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-06-06T14:09:29.159092","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Trained on low-res+high-res, epoch one of four.","metadata":{"id":"1ed267bc"}},{"cell_type":"markdown","source":"# config","metadata":{"id":"ad37275e"}},{"cell_type":"code","source":"DEBUG = False","metadata":{"executionInfo":{"elapsed":3,"status":"ok","timestamp":1720363921498,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"7c803349","execution":{"iopub.status.busy":"2024-07-18T20:32:43.727963Z","iopub.execute_input":"2024-07-18T20:32:43.728226Z","iopub.status.idle":"2024-07-18T20:32:43.742102Z","shell.execute_reply.started":"2024-07-18T20:32:43.728196Z","shell.execute_reply":"2024-07-18T20:32:43.741428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{"id":"e7fff7a9"}},{"cell_type":"code","source":"%%capture\nimport os, gc\nos.environ[\"KERAS_BACKEND\"] = \"tensorflow\"\n!pip install keras==2.15.0\nfrom tensorflow.keras.callbacks import EarlyStopping, ModelCheckpoint\nfrom tensorflow.keras.callbacks import LearningRateScheduler, ReduceLROnPlateau\nfrom tensorflow.keras import backend as K\nimport tensorflow as tf\nprint(tf.__version__)","metadata":{"execution":{"iopub.status.busy":"2024-07-18T20:32:43.753045Z","iopub.execute_input":"2024-07-18T20:32:43.753287Z","iopub.status.idle":"2024-07-18T20:33:07.170612Z","shell.execute_reply.started":"2024-07-18T20:32:43.753261Z","shell.execute_reply":"2024-07-18T20:33:07.16929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nfrom sklearn.model_selection import KFold\nimport sklearn\nimport matplotlib.pyplot as plt\nimport pickle\nimport math\nimport shutil","metadata":{"executionInfo":{"elapsed":6090,"status":"ok","timestamp":1720363930566,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"7482513c","execution":{"iopub.status.busy":"2024-07-18T20:33:07.172785Z","iopub.execute_input":"2024-07-18T20:33:07.173347Z","iopub.status.idle":"2024-07-18T20:33:09.179054Z","shell.execute_reply.started":"2024-07-18T20:33:07.173315Z","shell.execute_reply":"2024-07-18T20:33:09.178075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Save folder","metadata":{"id":"e20deb26"}},{"cell_type":"code","source":"save_folder = '/kaggle/working'\ntry:\n    os.mkdir(save_folder)\nexcept Exception as e:\n    print('exception error:')\n    print(e)","metadata":{"executionInfo":{"elapsed":21870,"status":"ok","timestamp":1720363952405,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"892a1746","outputId":"90a514d4-0940-4ef3-f607-8c92643b14f5","execution":{"iopub.status.busy":"2024-07-18T20:34:42.704118Z","iopub.execute_input":"2024-07-18T20:34:42.705212Z","iopub.status.idle":"2024-07-18T20:34:42.709855Z","shell.execute_reply.started":"2024-07-18T20:34:42.705174Z","shell.execute_reply":"2024-07-18T20:34:42.708992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Connect to TPU","metadata":{"id":"d38b7ce6"}},{"cell_type":"code","source":"# Configure Strategy. Assume TPU...if not set default for GPU\ntpu = None\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver()\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.TPUStrategy(tpu)\n    tpu_connceted = True\n    print(\"on TPU\")\n    print(\"REPLICAS: \", strategy.num_replicas_in_sync)\nexcept Exception as e:\n    print('exception error:')\n    print(e)\n    strategy = tf.distribute.get_strategy()\n\nprint(strategy)","metadata":{"executionInfo":{"elapsed":13421,"status":"ok","timestamp":1720363965819,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"8f94d825","outputId":"efdb1cd3-9ecc-43f3-c3ef-8fbe812b9b70","execution":{"iopub.status.busy":"2024-07-18T20:34:42.711478Z","iopub.execute_input":"2024-07-18T20:34:42.711743Z","iopub.status.idle":"2024-07-18T20:34:51.693223Z","shell.execute_reply.started":"2024-07-18T20:34:42.711717Z","shell.execute_reply":"2024-07-18T20:34:51.692241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TFRecords path","metadata":{"id":"1ddba393"}},{"cell_type":"code","source":"gcp_data = '/kaggle/input/leap-gcp-path-tfrecs-s'\nall_lens_lr_concat = pickle.load(open(f'{gcp_data}/all_lens_lr_concat.p', 'br'))\nall_tffiles_lr_concat = pickle.load(open(f'{gcp_data}/all_tffiles_lr_concat.p', 'br'))\nall_lens_hr_concat = pickle.load(open(f'{gcp_data}/all_lens_hr_concat.p', 'br'))\nall_tffiles_hr_concat = pickle.load(open(f'{gcp_data}/all_tffiles_hr_concat.p', 'br'))\n\ngcp_data = '/kaggle/input/leap-gcp-path-tfrecs-hr'\nall_lens_hr_concat_TOTAL = pickle.load(open(f'{gcp_data}/all_lens_hr_concat_TOTAL.p', 'br'))\nall_tffiles_hr_concat_TOTAL = pickle.load(open(f'{gcp_data}/all_tffiles_hr_concat_TOTAL.p', 'br'))","metadata":{"executionInfo":{"elapsed":2240,"status":"ok","timestamp":1720363968029,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"TPQOLApVXWa2","_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-07-18T20:34:51.694359Z","iopub.execute_input":"2024-07-18T20:34:51.694616Z","iopub.status.idle":"2024-07-18T20:34:52.014463Z","shell.execute_reply.started":"2024-07-18T20:34:51.694589Z","shell.execute_reply":"2024-07-18T20:34:52.013463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Auxiliary data","metadata":{}},{"cell_type":"code","source":"%%time\n\ninput_folder='/kaggle/working/input'\ntry:\n    os.mkdir(input_folder)\nexcept Exception as e:\n    print('exception error:')\n    print(e)","metadata":{"executionInfo":{"elapsed":7226,"status":"ok","timestamp":1720363975244,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"e8904671","outputId":"73051e90-e661-4ec5-9c80-d7b1675bbf69","execution":{"iopub.status.busy":"2024-07-18T20:34:52.015656Z","iopub.execute_input":"2024-07-18T20:34:52.01595Z","iopub.status.idle":"2024-07-18T20:34:52.021878Z","shell.execute_reply.started":"2024-07-18T20:34:52.015924Z","shell.execute_reply":"2024-07-18T20:34:52.021029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stds_path = '/kaggle/input/leap-sample-submission-stds-ds/stds.p'\nstds_new_path = '/kaggle/input/leap-sample-submission-stds-ds/stds_new.p'\n\nFEAT_COLS_path = '/kaggle/input/leap-sample-submission-stds-ds/FEAT_COLS.p'\nTARGET_COLS_path = '/kaggle/input/leap-sample-submission-stds-ds/TARGET_COLS.p'\ndata_path_MSM = '/kaggle/input/leap-msm-ds'\ndata_path_gdata ='/kaggle/input/leap-gdata-ds'\n\ny_max_per_col_all_path = '/kaggle/input/leap-sample-submission-stds-ds/y_max_per_col_all.p'\ny_min_per_col_all_path = '/kaggle/input/leap-sample-submission-stds-ds/y_min_per_col_all.p'","metadata":{"executionInfo":{"elapsed":327,"status":"ok","timestamp":1720363975568,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"8f57da20","execution":{"iopub.status.busy":"2024-07-18T20:34:52.023896Z","iopub.execute_input":"2024-07-18T20:34:52.024137Z","iopub.status.idle":"2024-07-18T20:34:52.034056Z","shell.execute_reply.started":"2024-07-18T20:34:52.024113Z","shell.execute_reply":"2024-07-18T20:34:52.033335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"[area, hyai, hyam, hybi, hybm, ilev, lat, lev, lon] = pickle.load(open(f'{data_path_gdata}/gdata.p', 'br'))\nlat = tf.convert_to_tensor(lat)\nlon = tf.convert_to_tensor(lon)\nlat_mean = tf.math.reduce_mean(lat)\nlat_std = tf.math.reduce_std(lat)\nlon_mean = tf.math.reduce_mean(lon)\nlon_std = tf.math.reduce_std(lon)\n\nlat = (lat - lat_mean)/lat_std\nlon = (lon - lon_mean)/lon_std\n\nlat = tf.cast(lat, tf.float32)\nlon = tf.cast(lon, tf.float32)\n\nday_time_sin = tf.convert_to_tensor(np.sin(np.asarray(range(12*31))*2*np.pi/(12*31)), dtype = tf.float32)\nday_time_cos = tf.convert_to_tensor(np.cos(np.asarray(range(12*31))*2*np.pi/(12*31)), dtype = tf.float32)\nsec_time_sin = tf.convert_to_tensor(np.sin(np.asarray(range(24*60*60))*2*np.pi/(24*60*60)), dtype = tf.float32)\nsec_time_cos = tf.convert_to_tensor(np.cos(np.asarray(range(24*60*60))*2*np.pi/(24*60*60)), dtype = tf.float32)","metadata":{"executionInfo":{"elapsed":28,"status":"ok","timestamp":1720363975568,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"4733671c","execution":{"iopub.status.busy":"2024-07-18T20:34:52.034907Z","iopub.execute_input":"2024-07-18T20:34:52.035122Z","iopub.status.idle":"2024-07-18T20:34:52.091918Z","shell.execute_reply.started":"2024-07-18T20:34:52.0351Z","shell.execute_reply":"2024-07-18T20:34:52.091131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"[area_hr, hyai_hr, hyam_hr, hybi_hr, hybm_hr, ilev_hr, lat_hr, lev_hr, lon_hr] = pickle.load(open(f'{data_path_gdata}/data_hr.p', 'br'))\nlat_hr = tf.convert_to_tensor(lat_hr)\nlon_hr = tf.convert_to_tensor(lon_hr)\nlat_mean_hr = tf.math.reduce_mean(lat_hr)\nlat_std_hr = tf.math.reduce_std(lat_hr)\n\nlat_hr = (lat_hr - lat_mean)/lat_std\nlon_hr = (lon_hr - lon_mean)/lon_std\n\nlat_hr = tf.cast(lat_hr, tf.float32)\nlon_hr = tf.cast(lon_hr, tf.float32)","metadata":{"executionInfo":{"elapsed":27,"status":"ok","timestamp":1720363975568,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"rDbasH48cOKD","execution":{"iopub.status.busy":"2024-07-18T20:34:52.092843Z","iopub.execute_input":"2024-07-18T20:34:52.093086Z","iopub.status.idle":"2024-07-18T20:34:52.106599Z","shell.execute_reply.started":"2024-07-18T20:34:52.093063Z","shell.execute_reply":"2024-07-18T20:34:52.105914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FEAT_COLS = pickle.load(open(FEAT_COLS_path, 'br'))\nTARGET_COLS = pickle.load(open(TARGET_COLS_path, 'br'))","metadata":{"executionInfo":{"elapsed":28,"status":"ok","timestamp":1720363975569,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"4086d2ac","execution":{"iopub.status.busy":"2024-07-18T20:34:52.107536Z","iopub.execute_input":"2024-07-18T20:34:52.107772Z","iopub.status.idle":"2024-07-18T20:34:52.118618Z","shell.execute_reply.started":"2024-07-18T20:34:52.107749Z","shell.execute_reply":"2024-07-18T20:34:52.117897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_cols = FEAT_COLS\ny_cols = TARGET_COLS\n\nX_cols_series_parents = ['state_t', 'state_q0001', 'state_q0002', 'state_q0003', 'state_u', 'state_v',\n                         'pbuf_ozone', 'pbuf_CH4', 'pbuf_N2O']\nX_cols_series = [x for x in X_cols if np.mean([x.startswith(y) for y in X_cols_series_parents])>0]\nX_cols_series_NOT = [x for x in X_cols if x not in X_cols_series]\nprint(len(X_cols_series)+len(X_cols_series_NOT) == len(X_cols))\nX_cols_series_indices = np.asarray([np.where(np.asarray(X_cols) == x)[0][0] for x in X_cols_series])\nX_cols_series_NOT_indices = np.asarray([np.where(np.asarray(X_cols) == x)[0][0] for x in X_cols_series_NOT])\n\ny_cols_series_parents = ['ptend_t', 'ptend_q0001', 'ptend_q0002', 'ptend_q0003', 'ptend_u', 'ptend_v']\ny_cols_series = [x for x in y_cols if np.mean([x.startswith(y) for y in y_cols_series_parents])>0]\ny_cols_series_NOT = [x for x in y_cols if x not in y_cols_series]\nprint(len(y_cols_series)+len(y_cols_series_NOT) == len(y_cols))\ny_cols_series_indices = np.asarray([np.where(np.asarray(y_cols) == x)[0][0] for x in y_cols_series])\ny_cols_series_NOT_indices = np.asarray([np.where(np.asarray(y_cols) == x)[0][0] for x in y_cols_series_NOT])","metadata":{"executionInfo":{"elapsed":27,"status":"ok","timestamp":1720363975569,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"4dea0970","outputId":"1c5d5800-b004-4f03-828f-3f2c9d8c16c5","execution":{"iopub.status.busy":"2024-07-18T20:34:52.119806Z","iopub.execute_input":"2024-07-18T20:34:52.120066Z","iopub.status.idle":"2024-07-18T20:34:52.157767Z","shell.execute_reply.started":"2024-07-18T20:34:52.12004Z","shell.execute_reply":"2024-07-18T20:34:52.156927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stds = pickle.load(open(stds_path, 'br'))\nstds_new = pickle.load(open(stds_new_path, 'br'))\nstds_new = stds_new.astype(np.float64)","metadata":{"executionInfo":{"elapsed":17,"status":"ok","timestamp":1720363975569,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"298912c5","execution":{"iopub.status.busy":"2024-07-18T20:34:52.158724Z","iopub.execute_input":"2024-07-18T20:34:52.158968Z","iopub.status.idle":"2024-07-18T20:34:52.168828Z","shell.execute_reply.started":"2024-07-18T20:34:52.158945Z","shell.execute_reply":"2024-07-18T20:34:52.168079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configs","metadata":{"id":"183343d7"}},{"cell_type":"code","source":"if DEBUG:\n    N_EPOCHS = 40\n    batch_size = 2\nelse:\n    N_EPOCHS = 40\n    batch_size = 2048\nval_batch_size = batch_size\nN_WARMUP_EPOCHS = 0\nLR_MAX = 1e-3\nLR_MIN = 8.4e-4\nWD_RATIO = 4.0\nWARMUP_METHOD = \"exp\"\nlr_factor = 0.7\nPAD = 0.0","metadata":{"executionInfo":{"elapsed":17,"status":"ok","timestamp":1720363975570,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"79591b91","execution":{"iopub.status.busy":"2024-07-18T20:34:52.169945Z","iopub.execute_input":"2024-07-18T20:34:52.170203Z","iopub.status.idle":"2024-07-18T20:34:52.174344Z","shell.execute_reply.started":"2024-07-18T20:34:52.170176Z","shell.execute_reply":"2024-07-18T20:34:52.173653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(all_lens_lr_concat))\nprint(len(all_lens_hr_concat))\nprint(len(all_lens_hr_concat_TOTAL))\nprint(np.sum(all_lens_lr_concat))\nprint(np.sum(all_lens_hr_concat_TOTAL))","metadata":{"executionInfo":{"elapsed":16,"status":"ok","timestamp":1720363975570,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"2pj2fDyR_eyL","outputId":"786ea7c8-b299-49c5-9810-fecf2e2637f5","execution":{"iopub.status.busy":"2024-07-18T20:34:52.175199Z","iopub.execute_input":"2024-07-18T20:34:52.175415Z","iopub.status.idle":"2024-07-18T20:34:52.184588Z","shell.execute_reply.started":"2024-07-18T20:34:52.175392Z","shell.execute_reply":"2024-07-18T20:34:52.183903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rand_seed = np.random.randint(10000)\nprint(rand_seed)\nrng = np.random.default_rng(rand_seed)\nrng.shuffle(all_tffiles_hr_concat_TOTAL[20:])","metadata":{"executionInfo":{"elapsed":10,"status":"ok","timestamp":1720363975571,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"ZGr_OdExNRnP","execution":{"iopub.status.busy":"2024-07-18T20:34:52.185488Z","iopub.execute_input":"2024-07-18T20:34:52.185852Z","iopub.status.idle":"2024-07-18T20:34:52.205927Z","shell.execute_reply.started":"2024-07-18T20:34:52.185822Z","shell.execute_reply":"2024-07-18T20:34:52.20513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    val_1_tffiles = all_tffiles_lr_concat[:1]\n    val_2_tffiles = all_tffiles_lr_concat[:1]\n    val_3_tffiles = all_tffiles_hr_concat[:1]\n    val_4_tffiles = all_tffiles_lr_concat[:1]\n    train_tffiles = all_tffiles_lr_concat[1:3]\n    lens_tffiles = np.concatenate([all_lens_lr_concat[1:2]], axis = 0)\nelse:\n    val_1_tffiles = all_tffiles_lr_concat[:20]\n    val_2_tffiles = all_tffiles_lr_concat[20:40]\n    val_3_tffiles = all_tffiles_hr_concat[0:20]\n    val_4_tffiles = all_tffiles_lr_concat[40:60]\n    train_tffiles = []\n    for i in range(11):\n        train_tffiles_temp = all_tffiles_lr_concat[40:].copy()\n        train_tffiles_temp_HR = all_tffiles_hr_concat_TOTAL[20+i*8060:20+(i+1)*8060].copy()\n        tffiles_temp_concat = np.concatenate([train_tffiles_temp, train_tffiles_temp_HR], axis = 0)\n        np.random.shuffle(tffiles_temp_concat)\n        train_tffiles.append(tffiles_temp_concat)\n    train_tffiles = np.concatenate(train_tffiles)\n    lens_tffiles = np.concatenate([all_lens_lr_concat[40:], all_lens_hr_concat_TOTAL[20:20+8060]], axis = 0)\n\nnum_val_1 = len(val_1_tffiles)*5000\nnum_val_2 = len(val_2_tffiles)*5000\nnum_val_3 = len(val_3_tffiles)*5000\nnum_val_4 = len(val_4_tffiles)*5000\nnum_train = np.sum(lens_tffiles)\nif DEBUG:\n    num_val_1 = 64\n    num_val_2 = 64\n    num_val_3 = 64\n    num_val_4 = 64\n    num_train = 64\nprint(f'num train: {num_train}')\n\nsteps_per_epoch = int(num_train/batch_size/10)\nif DEBUG:\n    steps_per_epoch = 2\n    \nval_steps_per_epoch = num_val_1//val_batch_size\nval_steps_per_epoch_2 = num_val_2//val_batch_size\nval_steps_per_epoch_3 = num_val_3//val_batch_size\nval_steps_per_epoch_4 = num_val_4//val_batch_size\n\nval_samples_num_1 = val_steps_per_epoch*val_batch_size\nval_samples_num_2 = val_steps_per_epoch_2*val_batch_size\nval_samples_num_3 = val_steps_per_epoch_3*val_batch_size\nval_samples_num_4 = val_steps_per_epoch_4*val_batch_size\nprint(f'steps_per_epoch: {steps_per_epoch}')\nprint(val_steps_per_epoch)\nprint(val_samples_num_1)\nprint(val_samples_num_2)\nprint(val_samples_num_3)\nprint(val_samples_num_4)\n\n\nepoch_samples_num = steps_per_epoch*batch_size\nif DEBUG:\n    epoch_samples_num = 10\nprint(epoch_samples_num)","metadata":{"executionInfo":{"elapsed":345,"status":"ok","timestamp":1720363975907,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"378b7ad4","outputId":"0d9350d0-40eb-4e61-b3fe-04cd2c7c4451","execution":{"iopub.status.busy":"2024-07-18T20:34:52.209324Z","iopub.execute_input":"2024-07-18T20:34:52.209615Z","iopub.status.idle":"2024-07-18T20:34:52.220185Z","shell.execute_reply.started":"2024-07-18T20:34:52.209587Z","shell.execute_reply":"2024-07-18T20:34:52.219325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"ceed","metadata":{"id":"-XNZUTeh-9WG"}},{"cell_type":"markdown","source":"# Auxiliary data process","metadata":{"id":"bdd8fcc5"}},{"cell_type":"code","source":"y_min_per_col_all = pickle.load((open(y_min_per_col_all_path, 'br')))\ny_max_per_col_all = pickle.load((open(y_max_per_col_all_path, 'br')))\n\nmin_y = pickle.load((open(f'{data_path_MSM}/min_y.p', 'br')))\nmax_y = pickle.load((open(f'{data_path_MSM}/max_y.p', 'br')))\nmean_y = pickle.load((open(f'{data_path_MSM}/mean_y.p', 'br')))\nstd_y = pickle.load((open(f'{data_path_MSM}/std_y.p', 'br')))\n\nmin_col_not = pickle.load((open(f'{data_path_MSM}/min_col_not.p', 'br')))\nmax_col_not = pickle.load((open(f'{data_path_MSM}/max_col_not.p', 'br')))\nmean_col_not = pickle.load((open(f'{data_path_MSM}/mean_col_not.p', 'br')))\nstd_col_not = pickle.load((open(f'{data_path_MSM}/std_col_not.p', 'br')))\n\nmin_col_not_test = pickle.load((open(f'{data_path_MSM}/min_col_not_test.p', 'br')))\nmax_col_not_test = pickle.load((open(f'{data_path_MSM}/max_col_not_test.p', 'br')))\nmean_col_not_test = pickle.load((open(f'{data_path_MSM}/mean_col_not_test.p', 'br')))\nstd_col_not_test = pickle.load((open(f'{data_path_MSM}/std_col_not_test.p', 'br')))\n\nx_col_min = pickle.load((open(f'{data_path_MSM}/x_col_min.p', 'br')))\nx_col_max = pickle.load((open(f'{data_path_MSM}/x_col_max.p', 'br')))\nx_col_mean = pickle.load((open(f'{data_path_MSM}/x_col_mean.p', 'br')))\nx_col_std = pickle.load((open(f'{data_path_MSM}/x_col_std.p', 'br')))\n\nx_col_min_test = pickle.load((open(f'{data_path_MSM}/x_col_min_test.p', 'br')))\nx_col_max_test = pickle.load((open(f'{data_path_MSM}/x_col_max_test.p', 'br')))\nx_col_mean_test = pickle.load((open(f'{data_path_MSM}/x_col_mean_test.p', 'br')))\nx_col_std_test = pickle.load((open(f'{data_path_MSM}/x_col_std_test.p', 'br')))\n\nX_total_min = pickle.load((open(f'{data_path_MSM}/X_total_min.p', 'br')))\nX_total_max = pickle.load((open(f'{data_path_MSM}/X_total_max.p', 'br')))\nX_total_mean = pickle.load((open(f'{data_path_MSM}/x_total_mean.p', 'br')))\nX_total_std = pickle.load((open(f'{data_path_MSM}/x_total_std.p', 'br')))\n\nX_total_min_test = pickle.load((open(f'{data_path_MSM}/X_total_min_test.p', 'br')))\nX_total_max_test = pickle.load((open(f'{data_path_MSM}/X_total_max_test.p', 'br')))\nX_total_mean_test = pickle.load((open(f'{data_path_MSM}/x_total_mean_test.p', 'br')))\nX_total_std_test = pickle.load((open(f'{data_path_MSM}/x_total_std_test.p', 'br')))\n\ny_total_min = pickle.load((open(f'{data_path_MSM}/y_total_min.p', 'br')))\ny_total_max = pickle.load((open(f'{data_path_MSM}/y_total_max.p', 'br')))\ny_total_mean = pickle.load((open(f'{data_path_MSM}/y_total_mean.p', 'br')))\ny_total_std = pickle.load((open(f'{data_path_MSM}/y_total_std.p', 'br')))","metadata":{"executionInfo":{"elapsed":24,"status":"ok","timestamp":1720363975907,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"806c8e7a","execution":{"iopub.status.busy":"2024-07-18T20:34:52.221288Z","iopub.execute_input":"2024-07-18T20:34:52.221569Z","iopub.status.idle":"2024-07-18T20:34:52.374665Z","shell.execute_reply.started":"2024-07-18T20:34:52.221544Z","shell.execute_reply":"2024-07-18T20:34:52.373806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stds_temp = np.where(std_y == 0, 0, 1/np.where(std_y>0, std_y, 1))\nstds = stds_temp[0]","metadata":{"executionInfo":{"elapsed":22,"status":"ok","timestamp":1720363975907,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"H8gPq7dQszQ8","outputId":"4480ad62-837f-4cfb-8a24-89f771905c8c","execution":{"iopub.status.busy":"2024-07-18T20:34:52.375651Z","iopub.execute_input":"2024-07-18T20:34:52.375936Z","iopub.status.idle":"2024-07-18T20:34:52.379821Z","shell.execute_reply.started":"2024-07-18T20:34:52.375909Z","shell.execute_reply":"2024-07-18T20:34:52.378968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean_col_not = mean_col_not[0]\nstd_col_not = std_col_not[0]\nmin_col_not_norm = ((np.where(min_col_not<min_col_not_test, min_col_not , min_col_not_test)-mean_col_not)/std_col_not)[0]\n\nX_total_mean = X_total_mean[0]\nX_total_std = X_total_std[0]\nX_total_norm_min = ((np.where(X_total_min<X_total_min_test, X_total_min , X_total_min_test)-X_total_mean)/X_total_std)[0]\n\nx_col_mean = x_col_mean[0]\nx_col_std = x_col_std[0]\nx_col_norm_min = ((np.where(x_col_min<x_col_min_test, x_col_min , x_col_min_test)-x_col_mean)/x_col_std)[0]\n\nmean_y = mean_y[0]\n\nx_col_min = x_col_min[0]","metadata":{"executionInfo":{"elapsed":15,"status":"ok","timestamp":1720363975908,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"dbffcd75","execution":{"iopub.status.busy":"2024-07-18T20:34:52.380937Z","iopub.execute_input":"2024-07-18T20:34:52.381214Z","iopub.status.idle":"2024-07-18T20:34:52.390715Z","shell.execute_reply.started":"2024-07-18T20:34:52.381184Z","shell.execute_reply":"2024-07-18T20:34:52.389871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"norm_y_min = tf.convert_to_tensor((y_min_per_col_all- mean_y)*stds)\nnorm_y_max = tf.convert_to_tensor((y_max_per_col_all- mean_y)*stds)","metadata":{"executionInfo":{"elapsed":11,"status":"ok","timestamp":1720363975908,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"vm_X_uh-V9Q9","outputId":"e35da13c-9c1c-49a5-e9c6-c854b0f59a16","execution":{"iopub.status.busy":"2024-07-18T20:34:52.391761Z","iopub.execute_input":"2024-07-18T20:34:52.392029Z","iopub.status.idle":"2024-07-18T20:34:52.401392Z","shell.execute_reply.started":"2024-07-18T20:34:52.392004Z","shell.execute_reply":"2024-07-18T20:34:52.400543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TF data pipeline","metadata":{}},{"cell_type":"code","source":"def decode_tfrec(record_bytes):\n    schema = {}\n    schema[\"x1\"] = tf.io.VarLenFeature(dtype=tf.string)\n    schema[\"x2\"] = tf.io.VarLenFeature(dtype=tf.string)\n    schema[\"x3\"] = tf.io.VarLenFeature(dtype=tf.string)\n    schema[\"x4\"] = tf.io.VarLenFeature(dtype=tf.string)\n    schema[\"time\"] = tf.io.VarLenFeature(dtype=tf.int64)\n    schema[\"gidx\"] = tf.io.VarLenFeature(dtype=tf.int64)\n    schema[\"idx\"] = tf.io.VarLenFeature(dtype=tf.int64)\n    schema[\"batch_index\"] = tf.io.VarLenFeature(dtype=tf.int64)\n    features = tf.io.parse_single_example(record_bytes, schema)\n\n    data_cols = tf.sparse.to_dense(features[\"x1\"])\n    data_cols = tf.io.parse_tensor(data_cols[0], out_type=tf.float64)\n    data_cols_not = tf.sparse.to_dense(features[\"x2\"])\n    data_cols_not = tf.io.parse_tensor(data_cols_not[0], out_type=tf.float64)\n    data_cols_targets = tf.sparse.to_dense(features[\"x3\"])\n    data_cols_targets = tf.io.parse_tensor(data_cols_targets[0], out_type=tf.float64)\n    data_cols_not_target = tf.sparse.to_dense(features[\"x4\"])\n    data_cols_not_target = tf.io.parse_tensor(data_cols_not_target[0], out_type=tf.float64)\n\n    time = tf.sparse.to_dense(features[\"time\"])\n    gidx = tf.sparse.to_dense(features[\"gidx\"])\n    gidx = gidx[0]\n    res = 0\n    if gidx<0:\n        gidx = -gidx-1\n        res = 1\n    idx = tf.sparse.to_dense(features[\"idx\"])\n    batch_index = tf.sparse.to_dense(features[\"batch_index\"])\n\n    out = {}\n    out['X_col']  = data_cols\n    out['X_col_not']  = data_cols_not\n    out['y']  = tf.concat([data_cols_targets, data_cols_not_target], axis = 0)\n    out['time']  = time\n    out['gidx']  = gidx\n    out['res'] = res\n    return out\n\ndef get_timespace(x, lat_hr, lon_hr, lat, lon, day_time_sin, day_time_cos, sec_time_sin, sec_time_cos):\n    res = x['res']\n    gidx = x['gidx']\n    time = x['time']\n    if res == 0:\n        x_lat = lat_hr[gidx]\n        x_lon = lon_hr[gidx]\n    else:\n        x_lat = lat[gidx]\n        x_lon = lon[gidx]\n\n    day_sin = day_time_sin[(time[1]-1)*31+time[2]-1]\n    day_cos = day_time_cos[(time[1]-1)*31+time[2]-1]\n    sec_sin = sec_time_sin[time[3]]\n    sec_cos = sec_time_cos[time[3]]\n\n    x.pop('gidx')\n    x.pop('time')\n\n    x['lat']  = x_lat\n    x['lon']  = x_lon\n    x['day_sin']  = day_sin\n    x['day_cos']  = day_cos\n    x['sec_sin']  = sec_sin\n    x['sec_cos']  = sec_cos\n    return x\n\ndef get_labels(x):\n    return x['y']\n\ndef get_labels_2(x):\n    return x['normed_y']\n\ndef get_normed(x, valid = False):\n    y_norm =  (x['y'] - mean_y)*stds\n    y_norm = tf.cast(y_norm, tf.float32)\n    \n    x_col_not_norm = (x['X_col_not'] - mean_col_not)/std_col_not\n\n    x_total_norm = tf.reshape(x['X_col'], [-1, 9, 60])\n    x_total_norm = (x_total_norm - X_total_mean)/X_total_std\n    x_total_norm = tf.reshape(x_total_norm, [60*x_total_norm.shape[1]])\n\n\n    x_col_norm = (x['X_col'] - x_col_mean)/x_col_std\n    x_col_norm_log = tf.where((x_col_norm-x_col_norm_min+1)>=1, tf.math.log(x_col_norm-x_col_norm_min+1),\n                                    -tf.math.log(1+1-(x_col_norm-x_col_norm_min+1)))\n\n\n    wind = tf.reshape(x['X_col'], [-1, 9, 60])\n    wind = wind[:, 4:6, :]\n    wind = (tf.math.reduce_sum(wind**2, axis = 1))**0.5\n    wind = wind[0]\n    wind = ((wind - tf.math.reduce_mean((tf.reshape(x_col_mean, [9, 60])[4:6]), axis = 0))/\n                    tf.math.reduce_sum((tf.reshape(x_col_std, [9, 60])[4:6]), axis = 0))\n\n\n    cutoff = 30\n    square_cutoff = cutoff**0.5\n    x_col_norm = tf.where(x_col_norm>cutoff, x_col_norm**0.5+cutoff-square_cutoff, x_col_norm)\n    x_col_norm = tf.where(x_col_norm<-cutoff, -tf.math.abs(x_col_norm)**0.5-cutoff+square_cutoff, x_col_norm)\n\n    wind = tf.where(wind>cutoff, wind**0.5+cutoff-square_cutoff, wind)\n    wind = tf.where(wind<-cutoff, -tf.math.abs(wind)**0.5-cutoff+square_cutoff, wind)\n\n\n    x_col_not_norm = tf.cast(x_col_not_norm, tf.float32)\n    x_total_norm = tf.cast(x_total_norm, tf.float32)\n    x_col_norm = tf.cast(x_col_norm, tf.float32)\n    x_col_norm_log = tf.cast(x_col_norm_log, tf.float32)\n    wind = tf.cast(wind, tf.float32)\n\n    cutoff_2 = 86.0\n    log_cutoff = tf.math.log(cutoff_2)\n    x_col_norm = tf.where(x_col_norm>cutoff_2, tf.math.log(x_col_norm)+cutoff_2-log_cutoff, x_col_norm)\n    x_col_norm = tf.where(x_col_norm<-cutoff_2, -tf.math.log(-x_col_norm)-cutoff_2+log_cutoff, x_col_norm)\n\n    x_col_not_norm = tf.where(x_col_not_norm>cutoff_2, tf.math.log(x_col_not_norm)+cutoff_2-log_cutoff, x_col_not_norm)\n    x_col_not_norm = tf.where(x_col_not_norm<-cutoff_2, -tf.math.log(-x_col_not_norm)-cutoff_2+log_cutoff, x_col_not_norm)\n\n    x_total_norm = tf.where(x_total_norm>cutoff_2, tf.math.log(x_total_norm)+cutoff_2-log_cutoff, x_total_norm)\n    x_total_norm = tf.where(x_total_norm<-cutoff_2, -tf.math.log(-x_total_norm)-cutoff_2+log_cutoff, x_total_norm)\n\n    wind = tf.where(wind>cutoff_2, tf.math.log(wind)+cutoff_2-log_cutoff, wind)\n    wind = tf.where(wind<-cutoff_2, -tf.math.log(-wind)-cutoff_2+log_cutoff, wind)\n\n    x.pop('X_col')\n    x.pop('X_col_not')\n    x.pop('y')\n\n    x['normed_y'] = y_norm\n    x['x_col_not_norm'] = x_col_not_norm\n    x['x_total_norm'] = x_total_norm\n    x['x_col_norm'] = x_col_norm\n    x['x_col_norm_log'] = x_col_norm_log\n    x['wind'] = wind\n    x['epoch'] = 0\n    return x\n\ndef get_concats(x):\n    x['concat_x'] = tf.concat([x['x_total_norm'], x['x_col_norm'], x['x_col_norm_log'],x['wind']], axis = 0)\n    x['concat_x_col_not'] = tf.concat([x['x_col_not_norm']], axis = 0)\n    x['concat_target'] = tf.concat([x['normed_y'], x['lat'][None], x['lon'][None],\n                                    x['day_sin'][None], x['day_cos'][None], x['sec_sin'][None], x['sec_cos'][None]], axis = 0)\n\n    x.pop('x_total_norm')\n    x.pop('x_col_norm')\n    x.pop('x_col_norm_log')\n    x.pop('x_col_not_norm')\n    x.pop('lat')\n    x.pop('lon')\n    x.pop('day_sin')\n    x.pop('day_cos')\n    x.pop('sec_sin')\n    x.pop('sec_cos')\n    x.pop('wind')\n    return x\n\ndef get_cols(x):\n    return x['concat_x']\n\ndef get_cols_not(x):\n    return x['concat_x_col_not']\n\ndef get_stateq002(x):\n    return x['X_col'][120:120+30]\n\ndef get_final_output(x):\n    x_concat = tf.concat([x['concat_x'], x['concat_x_col_not']], axis = 0)\n    concat_target = x['concat_target']\n    return x_concat, concat_target","metadata":{"executionInfo":{"elapsed":7,"status":"ok","timestamp":1720363975909,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"e37ffe84","execution":{"iopub.status.busy":"2024-07-18T20:34:52.402385Z","iopub.execute_input":"2024-07-18T20:34:52.40267Z","iopub.status.idle":"2024-07-18T20:34:52.431232Z","shell.execute_reply.started":"2024-07-18T20:34:52.402643Z","shell.execute_reply":"2024-07-18T20:34:52.430358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = tf.data.TFRecordDataset(\n    val_1_tffiles, num_parallel_reads=tf.data.AUTOTUNE, compression_type = 'GZIP').prefetch(tf.data.AUTOTUNE)\nds = ds.map(decode_tfrec, tf.data.AUTOTUNE)\nds = ds.map(lambda x: get_timespace(x, lat_hr, lon_hr, lat, lon, day_time_sin, day_time_cos, sec_time_sin, sec_time_cos), tf.data.AUTOTUNE)\nds = ds.map(get_normed, tf.data.AUTOTUNE)\nds = ds.map(get_concats, tf.data.AUTOTUNE)\n\nds_col = ds.map(get_cols, tf.data.AUTOTUNE)\nbatch = next(iter(ds_col))\ncol_len = batch.shape[0]\nprint(col_len)\n\nds_col_not = ds.map(get_cols_not, tf.data.AUTOTUNE)\nbatch = next(iter(ds_col_not))\ncol_not_len = batch.shape[0]\nprint(col_not_len)\n\nds = ds.map(get_final_output, tf.data.AUTOTUNE)\nbatch = next(iter(ds))\nprint(batch[0].shape, batch[1].shape)\ncol_num_y = batch[1].shape[0]\nprint(col_num_y)","metadata":{"executionInfo":{"elapsed":1745,"status":"ok","timestamp":1720363980041,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"1b81c163","outputId":"ded975f7-a53b-415d-df34-e9972cd79ba5","execution":{"iopub.status.busy":"2024-07-18T20:34:52.432268Z","iopub.execute_input":"2024-07-18T20:34:52.432586Z","iopub.status.idle":"2024-07-18T20:34:55.81077Z","shell.execute_reply.started":"2024-07-18T20:34:52.432558Z","shell.execute_reply":"2024-07-18T20:34:55.809618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_dataset(tffiles, batch_size, shuffle, split = 'valid', to_repeat = False, drop_reminder = True, cache = False, valid = False):\n    ds = tf.data.TFRecordDataset(\n        tffiles, num_parallel_reads=tf.data.AUTOTUNE, compression_type = 'GZIP').prefetch(tf.data.AUTOTUNE)\n    ds = ds.map(decode_tfrec, tf.data.AUTOTUNE)\n    ds = ds.map(lambda x: get_timespace(x, lat_hr, lon_hr, lat, lon, day_time_sin, day_time_cos, sec_time_sin, sec_time_cos), tf.data.AUTOTUNE)\n\n\n    if DEBUG:\n        ds = ds.take(64)\n\n    if cache:\n        ds = ds.cache()\n        samples_num = ds.reduce(0, lambda x,_: x+1).numpy()\n\n    if shuffle:\n        ds = ds.shuffle(shuffle, reshuffle_each_iteration = True)\n    if to_repeat:\n        ds = ds.repeat()\n\n    ds_label = ds.map(get_labels, tf.data.AUTOTUNE)\n    ds_q002 = ds.map(get_stateq002, tf.data.AUTOTUNE)\n\n    ds = ds.map(lambda x: get_normed(x, valid), tf.data.AUTOTUNE)\n    ds = ds.map(get_concats, tf.data.AUTOTUNE)\n\n    ds_label_2 = ds.map(get_labels_2, tf.data.AUTOTUNE)\n\n    ds = ds.map(get_final_output, tf.data.AUTOTUNE)\n\n    ds = ds.padded_batch(\n                batch_size, padding_values=(PAD, PAD),\n        padded_shapes=([col_len+col_not_len],[col_num_y]), drop_remainder=drop_reminder)\n\n    ds_label = ds_label.batch(batch_size)\n    ds_label_2 = ds_label_2.batch(batch_size)\n    ds_q002 = ds_q002.batch(batch_size)\n\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n    return ds, ds_label, ds_label_2, ds_q002","metadata":{"executionInfo":{"elapsed":6,"status":"ok","timestamp":1720363980042,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"48255e7b","execution":{"iopub.status.busy":"2024-07-18T20:34:55.812118Z","iopub.execute_input":"2024-07-18T20:34:55.812554Z","iopub.status.idle":"2024-07-18T20:34:55.821808Z","shell.execute_reply.started":"2024-07-18T20:34:55.812521Z","shell.execute_reply":"2024-07-18T20:34:55.820993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_1_ds, val_1_labels_ds, val_1_labels_ds_2, val_1_ds_q0002 = get_dataset(\n    val_1_tffiles, val_batch_size, shuffle = False, cache = True)","metadata":{"executionInfo":{"elapsed":23152,"status":"ok","timestamp":1720364003189,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"2a29beef","execution":{"iopub.status.busy":"2024-07-18T20:34:55.822823Z","iopub.execute_input":"2024-07-18T20:34:55.823129Z","iopub.status.idle":"2024-07-18T20:35:00.052764Z","shell.execute_reply.started":"2024-07-18T20:34:55.8231Z","shell.execute_reply":"2024-07-18T20:35:00.051758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch = next(iter(val_1_ds))\nbatch[0].shape ,batch[1].shape","metadata":{"executionInfo":{"elapsed":588,"status":"ok","timestamp":1720364003771,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"ef3ded02","outputId":"d3148a0e-fd0e-40c3-f7ff-e645a21fa80d","execution":{"iopub.status.busy":"2024-07-18T20:35:00.054119Z","iopub.execute_input":"2024-07-18T20:35:00.054405Z","iopub.status.idle":"2024-07-18T20:35:00.237427Z","shell.execute_reply.started":"2024-07-18T20:35:00.054374Z","shell.execute_reply":"2024-07-18T20:35:00.236636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Layers","metadata":{"id":"41079478"}},{"cell_type":"code","source":"class GLU(tf.keras.layers.Layer):\n    def __init__(self, **kwargs):\n        super().__init__(**kwargs)\n    def call(self, x, mask=None):\n        x,gate = tf.split(x, 2, axis = -1)\n        x = x*tf.keras.activations.swish(gate)\n        return x\n\nclass GLUMlp(tf.keras.layers.Layer):\n    def __init__(self, dim_expand, dim, **kwargs):\n        super().__init__(**kwargs)\n        self.dim_expand = dim_expand\n        self.dim = dim\n        self.dense_1 = tf.keras.layers.EinsumDense(\"abc,cd->abd\",output_shape=(None, self.dim_expand), activation = 'linear', bias_axes = 'd')\n        self.glu_1 = GLU()\n        self.dense_2 = tf.keras.layers.EinsumDense(\"abc,cd->abd\",output_shape=(None, self.dim), activation = 'linear', bias_axes = 'd')\n    def call(self, x, training = False):\n        x = self.dense_1(x)\n        x = self.glu_1(x)\n        x = self.dense_2(x)\n        return x\n\nclass ScaleBias(tf.keras.layers.Layer):\n    def __init__(self, **kwargs):\n        super().__init__(**kwargs)\n    def build(self, input_shape):\n        self.scale_bias = tf.keras.layers.EinsumDense(\"abc,c->abc\",output_shape=(None, input_shape[-1]),activation = 'linear', bias_axes = 'c')\n    def call(self, x, mask=None):\n        return self.scale_bias(x)\n\nclass TransformerEncoder(tf.keras.layers.Layer):\n    def __init__(self, embed_dim, num_heads, feed_forward_dim):\n        super().__init__()\n        self.att = tf.keras.layers.MultiHeadAttention(num_heads=num_heads, key_dim=embed_dim//num_heads)\n        self.ffn = GLUMlp(feed_forward_dim, embed_dim)\n        self.layer_norm_1 = tf.keras.layers.LayerNormalization(epsilon=1e-6)\n        self.layer_norm_2 = tf.keras.layers.LayerNormalization(epsilon=1e-6)\n        self.scale_bias_1 = ScaleBias()\n        self.scale_bias_2 = ScaleBias()\n    def call(self, x, training = None):\n        residual = x\n        x = self.att(x, x)\n        x = self.scale_bias_1(x)\n        x = self.layer_norm_1(x + residual)\n        residual = x\n        x = self.ffn(x, training = training)\n        x = self.scale_bias_2(x)\n        x = self.layer_norm_2(x + residual)\n        return x\n\n\nclass ECA(tf.keras.layers.Layer):\n    # TF implementation from https://www.kaggle.com/code/hoyso48/1st-place-solution-training\n    def __init__(self, kernel_size=5, **kwargs):\n        super().__init__(**kwargs)\n        self.kernel_size = kernel_size\n        self.conv = tf.keras.layers.Conv1D(1, kernel_size=kernel_size, strides=1, padding=\"same\", use_bias=False)\n    def call(self, inputs):\n        nn = tf.keras.layers.GlobalAveragePooling1D()(inputs)\n        nn = tf.expand_dims(nn, -1)\n        nn = self.conv(nn)\n        nn = tf.squeeze(nn, -1)\n        nn = tf.nn.sigmoid(nn)\n        nn = nn[:,None,:]\n        return inputs * nn\n\n\nclass HeadDense(tf.keras.layers.Layer):\n    def __init__(self, head_dim, **kwargs):\n        super().__init__(**kwargs)\n        self.head_dim = head_dim\n    def build(self, input_shape):\n        self.length = input_shape[1]\n        self.dim = input_shape[2]\n        self.dense = tf.keras.layers.EinsumDense(\"abc,cd->abd\",output_shape=(self.length, self.head_dim), activation = 'swish', bias_axes = 'd')\n    def call(self, x):\n        x = self.dense(x)\n        return x\n\nclass Conv1DBlockSqueezeformer(tf.keras.layers.Layer):\n    def __init__(self, channel_size, kernel_size, dilation_rate=1,\n                 expand_ratio=4, se_ratio=0.25, activation='swish', name=None, **kwargs):\n        super().__init__()\n        self.channel_size = channel_size\n        self.kernel_size = kernel_size\n        self.dilation_rate = dilation_rate\n        self.expand_ratio = expand_ratio\n        self.se_ratio = se_ratio\n        self.activation = activation\n        self.scale_bias = ScaleBias()\n        self.glu_layer = GLU()\n        self.ffn = GLUMlp(channel_size*4, channel_size)\n        self.layer_norm_2 = tf.keras.layers.LayerNormalization(epsilon=1e-6)\n        self.scale_bias_1 = ScaleBias()\n        self.scale_bias_2 = ScaleBias()\n    def build(self, input_shape):\n        self.length = input_shape[1]\n        self.channels_in = input_shape[2]\n        self.channels_expand = self.channels_in * self.expand_ratio\n        self.dwconv = tf.keras.layers.DepthwiseConv1D(self.kernel_size,dilation_rate=self.dilation_rate,padding='same',use_bias=False)\n        self.BatchNormalization_layer = tf.keras.layers.BatchNormalization(momentum=0.95)\n        self.conv_activation = tf.keras.layers.Activation(self.activation)\n        self.layernorm = tf.keras.layers.LayerNormalization(epsilon=1e-6)\n        self.ECA_layer = ECA()\n        self.expand = tf.keras.layers.EinsumDense(\"abc,cd->abd\",output_shape=(self.length, self.channels_expand), activation = 'linear', bias_axes = 'd')\n        self.project =tf.keras.layers.EinsumDense(\"abc,cd->abd\",output_shape=(self.length, self.channel_size), activation = 'linear', bias_axes = 'd')\n    def call(self, x, training = None):\n        skip = x\n        x = self.expand(x)\n        x = self.glu_layer(x)\n        x = self.dwconv(x)\n        x = self.BatchNormalization_layer(x)\n        x = self.conv_activation(x)\n        x = self.ECA_layer(x)\n        x = self.project(x)\n        x = self.scale_bias_1(x)\n\n        x = x+skip\n\n        residual = x\n        x = self.ffn(x)\n        x = self.scale_bias_2(x)\n        x = self.layer_norm_2(x + residual)\n        return x\n\nclass Reshape1(tf.keras.layers.Layer):\n    def __init__(self, col_len, **kwargs):\n        super().__init__(**kwargs)\n        self.col_len = col_len\n    def call(self, x):\n        x_seq = x[:, :self.col_len]\n        x_seq = tf.reshape(x_seq, [-1, tf.cast(self.col_len/60, tf.int32), 60])\n        x_seq = tf.transpose(x_seq, [0, 2, 1])\n        x_seq_N = x[:, self.col_len:]\n\n        x_seq_N = tf.expand_dims(x_seq_N, axis = 1)\n        x_seq_N = tf.repeat(x_seq_N, 60, axis = 1)\n\n        x = tf.concat([x_seq, x_seq_N], axis = -1)\n        return x\n\nclass Reshape2(tf.keras.layers.Layer):\n    def __init__(self, **kwargs):\n        super().__init__(**kwargs)\n    def call(self, x_pred, x_confidence):\n        x = x_pred\n        x_seq = x[:, :, :6]\n        x_seq = tf.transpose(x_seq, [0,2,1])\n        x_seq = tf.reshape(x_seq, [-1, 60*6])\n        x_seq_N = x[:, :, 6:]\n        x_seq_N = tf.reduce_mean(x_seq_N, axis = 1)\n        x1 = tf.concat([x_seq, x_seq_N], axis = -1)\n\n        x = x_confidence\n        x_seq = x[:, :, :6]\n        x_seq = tf.transpose(x_seq, [0,2,1])\n        x_seq = tf.reshape(x_seq, [-1, 60*6])\n        x_seq_N = x[:, :, 6:]\n        x_seq_N = tf.reduce_mean(x_seq_N, axis = 1)\n        x2 = tf.concat([x_seq, x_seq_N], axis = -1)\n\n        x = tf.concat([x1, x2], axis = -1)\n        return x","metadata":{"executionInfo":{"elapsed":7,"status":"ok","timestamp":1720364003771,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"4f63c39a","execution":{"iopub.status.busy":"2024-07-18T20:35:00.238869Z","iopub.execute_input":"2024-07-18T20:35:00.239168Z","iopub.status.idle":"2024-07-18T20:35:00.270909Z","shell.execute_reply.started":"2024-07-18T20:35:00.239139Z","shell.execute_reply":"2024-07-18T20:35:00.27003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss function","metadata":{}},{"cell_type":"code","source":"loss_mask = (stds_new!=0)[None]\nloss_mask = tf.cast(loss_mask, tf.float32)\ntemp = np.ones((loss_mask.shape))\ntemp[0, 120+12:120+26] = 0\nloss_mask = loss_mask*temp\nloss_mask = tf.pad(loss_mask, [[0,0], [0,col_num_y - len(stds_new)]], constant_values = 1)","metadata":{"execution":{"iopub.status.busy":"2024-07-18T20:35:00.271852Z","iopub.execute_input":"2024-07-18T20:35:00.272141Z","iopub.status.idle":"2024-07-18T20:35:00.296227Z","shell.execute_reply.started":"2024-07-18T20:35:00.272114Z","shell.execute_reply":"2024-07-18T20:35:00.295427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def loss_fn(x, preds):\n    confidence = preds[:, col_num_y:]\n    preds = preds[:, :col_num_y]\n    loss = tf.math.abs(x-preds)\n    loss = loss*loss_mask\n    loss_2 = tf.math.abs(loss-confidence)\n    loss_2 = loss_2*loss_mask\n    return tf.reduce_mean(loss+loss_2)\n\ndef metric_mae(x, preds):\n    x = x[:, :368]\n    preds = preds[:, :368]\n    loss = tf.math.abs(x-preds)\n    loss = loss*loss_mask[:, :368]\n    return tf.reduce_mean(loss)","metadata":{"executionInfo":{"elapsed":5,"status":"ok","timestamp":1720364003771,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"3dc3f687","execution":{"iopub.status.busy":"2024-07-18T20:35:00.297123Z","iopub.execute_input":"2024-07-18T20:35:00.297375Z","iopub.status.idle":"2024-07-18T20:35:00.322439Z","shell.execute_reply.started":"2024-07-18T20:35:00.297349Z","shell.execute_reply":"2024-07-18T20:35:00.32166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"def get_model(dim = 384, head_dim = 2048):\n    with strategy.scope():\n        inp1 = tf.keras.Input([col_len+col_not_len])\n        x = inp1\n\n        x = Reshape1(col_len)(x)\n        x = tf.keras.layers.Dense(dim, use_bias=False)(x)\n        x = tf.keras.layers.LayerNormalization(epsilon=1e-6)(x)\n\n        groups = 1\n        conv_filter = 15\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n        x = Conv1DBlockSqueezeformer(dim,conv_filter)(x)\n        x = TransformerEncoder(dim, 4, dim*4)(x)\n\n        x = HeadDense(head_dim)(x)\n        x = GLUMlp(head_dim*2, head_dim)(x)\n\n        x_pred = tf.keras.layers.Dense(20)(x)\n        x_confidence = tf.keras.layers.Dense(20)(x)\n\n        x = Reshape2()(x_pred, x_confidence)\n\n        model = tf.keras.Model(inp1, x)\n        return model","metadata":{"executionInfo":{"elapsed":7,"status":"ok","timestamp":1720364004416,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"84a4cae0","execution":{"iopub.status.busy":"2024-07-18T20:35:00.323361Z","iopub.execute_input":"2024-07-18T20:35:00.323605Z","iopub.status.idle":"2024-07-18T20:35:00.347425Z","shell.execute_reply.started":"2024-07-18T20:35:00.323579Z","shell.execute_reply":"2024-07-18T20:35:00.346682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    model = get_model(4, head_dim = 4)\nelse:\n    model = get_model()\n\nwith strategy.scope():\n    loss = loss_fn\n    optimizer = tf.keras.optimizers.AdamW(learning_rate=0.0005)\n    if DEBUG:\n        model.compile(loss=loss, optimizer=optimizer)\n    else:\n        model.compile(loss=loss, optimizer=optimizer, metrics = [metric_mae], steps_per_execution = 100)\n\nmodel(batch[0])\nmodel.summary()","metadata":{"executionInfo":{"elapsed":47377,"status":"ok","timestamp":1720364051787,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"203ecce5","outputId":"37e4c54f-ec63-4aed-a58a-5c1d404adbdb","execution":{"iopub.status.busy":"2024-07-18T20:35:00.34829Z","iopub.execute_input":"2024-07-18T20:35:00.348514Z","iopub.status.idle":"2024-07-18T20:35:15.987732Z","shell.execute_reply.started":"2024-07-18T20:35:00.34849Z","shell.execute_reply":"2024-07-18T20:35:15.986916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"\n# Validation","metadata":{"id":"7b55f2eb"}},{"cell_type":"code","source":"val_1_labels = [x.numpy() for x in val_1_labels_ds]\nval_1_labels = np.concatenate(val_1_labels, axis = 0)\nval_1_labels = val_1_labels[:val_samples_num_1]\n\nval_1_labels_2 = [x.numpy() for x in val_1_labels_ds_2]\nval_1_labels_2 = np.concatenate(val_1_labels_2, axis = 0)\nval_1_labels_2 = val_1_labels_2[:val_samples_num_1]\n\nval_1_q0002 = [x.numpy() for x in val_1_ds_q0002]\nval_1_q0002 = np.concatenate(val_1_q0002, axis = 0)\nval_1_q0002 = val_1_q0002[:val_samples_num_1]","metadata":{"executionInfo":{"elapsed":24880,"status":"ok","timestamp":1720364077510,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"8dde25b6","execution":{"iopub.status.busy":"2024-07-18T20:35:16.569085Z","iopub.execute_input":"2024-07-18T20:35:16.569387Z","iopub.status.idle":"2024-07-18T20:35:17.674741Z","shell.execute_reply.started":"2024-07-18T20:35:16.569357Z","shell.execute_reply":"2024-07-18T20:35:17.673622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_metric_list = []\n\njump = 2\nif DEBUG:\n    jump = 1\n\ndef valid_sub_func(model, val_dataset, val_labels, val_labels_valid2, val_q0002, to_save = False, idx = -1, epoch = -1):\n    val_labels = val_labels.astype(np.float64)\n\n    preds = model.predict(val_dataset, verbose = 0, batch_size = batch_size, steps = val_steps_per_epoch)\n    preds = preds[:, :368]\n    if to_save:\n        pickle.dump(preds, open(f'{save_folder_val}/preds_{idx}_{epoch}.p', 'bw'))\n    preds = preds.astype(np.float64)\n\n    metrics = np.asarray([sklearn.metrics.r2_score(val_labels[:, i], preds[:, i]) for i in range(368)])\n\n    for i in range(len(metrics)):\n        if metrics[i]<0:\n            preds[:,i] = 0\n\n    #print('Metric: ------------------')\n    metric = sklearn.metrics.r2_score(val_labels, preds)\n    #print(metric)\n\n    preds = preds + mean_y.reshape(1,-1)*stds\n    preds[:, np.where(stds_new == 0)] = 0\n    preds = preds/np.where(stds>0, stds, 1)\n    val_labels_valid2 = np.where(stds_new == 0, 0, val_labels_valid2)\n\n    #print('Metric: ------------------')\n    metric = sklearn.metrics.r2_score(val_labels_valid2,preds)\n    #print(metric)\n\n    preds[:, 120+12:120+28] = (-val_q0002[:, 12:28]/1200)\n\n    #print('Metric: ------------------')\n    metric = sklearn.metrics.r2_score(val_labels_valid2,preds)\n    print(metric)\n    #print('***')\n    return metrics","metadata":{"executionInfo":{"elapsed":4,"status":"ok","timestamp":1720364078915,"user":{"displayName":"shlomo ron","userId":"09718489046984556743"},"user_tz":-180},"id":"6558ed4b","execution":{"iopub.status.busy":"2024-07-18T20:35:17.68655Z","iopub.execute_input":"2024-07-18T20:35:17.687158Z","iopub.status.idle":"2024-07-18T20:35:17.701448Z","shell.execute_reply.started":"2024-07-18T20:35:17.687125Z","shell.execute_reply":"2024-07-18T20:35:17.700718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Validation inference","metadata":{"id":"d79bb5c7"}},{"cell_type":"code","source":"K.clear_session()\nif DEBUG:\n    model = get_model(4, head_dim = 4)\n    #model = get_model()\nelse:\n    model = get_model()","metadata":{"id":"8d03b64d","execution":{"iopub.status.busy":"2024-07-18T20:35:17.716394Z","iopub.execute_input":"2024-07-18T20:35:17.716642Z","iopub.status.idle":"2024-07-18T20:39:21.97586Z","shell.execute_reply.started":"2024-07-18T20:35:17.716618Z","shell.execute_reply":"2024-07-18T20:39:21.974769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_weights(f'/kaggle/input/leap-model-weights-lr/0p0-9epochs-E90.weights.h5')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics_1 = valid_sub_func(model, val_1_ds, val_1_labels_2, val_1_labels, val_1_q0002)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del val_1_labels, val_1_q0002\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test inference","metadata":{}},{"cell_type":"code","source":"def arrange_x_y(X_col, X_col_not,y):\n    x = {}\n    x['X_col'] = X_col\n    x['X_col_not'] = X_col_not\n    x['y'] = y\n    return x\n\ndef get_extra_vars_for_test_ds(x):\n    x['lat']  = 0.0\n    x['lon']  = 0.0\n    x['day_sin']  = 0.0\n    x['day_cos']  = 0.0\n    x['sec_sin']  = 0.0\n    x['sec_cos']  = 0.0\n    return x","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_dataset(X_train_col, X_train_col_not,y, batch_size, shuffle, split = 'valid', to_repeat = False, drop_reminder = True):\n    ds = tf.data.Dataset.from_tensor_slices((X_train_col, X_train_col_not,y))\n    ds = ds.map(arrange_x_y, tf.data.AUTOTUNE)\n    ds = ds.map(get_extra_vars_for_test_ds, tf.data.AUTOTUNE)\n\n    if DEBUG:\n        ds = ds.take(64)\n\n    ds = ds.cache()\n    samples_num = ds.reduce(0, lambda x,_: x+1).numpy()\n\n    if shuffle:\n        ds = ds.shuffle(shuffle, reshuffle_each_iteration = True)\n    if to_repeat:\n        ds = ds.repeat()\n\n    ds_label = ds.map(get_labels, tf.data.AUTOTUNE)\n    ds_q002 = ds.map(get_stateq002, tf.data.AUTOTUNE)\n    ds = ds.map(get_normed, tf.data.AUTOTUNE)\n    ds = ds.map(get_concats, tf.data.AUTOTUNE)\n    ds_label_2 = ds.map(get_labels_2, tf.data.AUTOTUNE)\n\n    ds = ds.map(get_final_output, tf.data.AUTOTUNE)\n\n    ds = ds.padded_batch(\n                batch_size, padding_values=(PAD, PAD),\n        padded_shapes=([col_len+col_not_len],[col_num_y]), drop_remainder=drop_reminder)\n\n    ds_label = ds_label.batch(batch_size)\n    ds_label_2 = ds_label_2.batch(batch_size)\n    ds_q002 = ds_q002.batch(batch_size)\n\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n    return ds, ds_label, ds_label_2, ds_q002","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import polars as pl","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_test_ds():\n    test_df = pl.read_csv('/kaggle/input/leap-atmospheric-physics-ai-climsim/test.csv')\n\n    X_train_col_not = test_df.select(X_cols_series_NOT).to_numpy()\n    X_train_col = test_df.select(X_cols_series).to_numpy()\n    y = X_train_col[:, :368]\n\n    test_ds, test_labels_ds, val_1_labels_ds_2, val_1_ds_q0002 = get_dataset(\n        X_train_col, X_train_col_not,y, val_batch_size, shuffle = False, drop_reminder = False)\n    del y, X_train_col, X_train_col_not, test_df\n    gc.collect()\n    return test_ds, test_labels_ds, val_1_labels_ds_2, val_1_ds_q0002\n\ntest_ds, test_labels_ds, val_1_labels_ds_2, val_1_ds_q0002 = get_test_ds()\nmetrics_best_epoch = metrics_1","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_q0002 = [x.numpy() for x in val_1_ds_q0002]\ntest_q0002 =np.concatenate(test_q0002, axis = 0)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"shutil.rmtree('/kaggle/working/input')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_preds(model, epoch):\n    preds = model.predict(test_ds, verbose = 0)\n    pickle.dump(preds, open(f'raw_preds_{epoch}.p', 'bw'))\n    preds = preds[:, :368]\n    preds = preds.astype(np.float64)\n\n    for i in range(len(metrics_best_epoch)):\n        if metrics_best_epoch[i]<0:\n            preds[:,i] = 0\n\n    preds = preds + mean_y.reshape(1,-1)*stds\n    preds[:, np.where(stds_new == 0)] = 0\n    preds = preds/np.where(stds>0, stds, 1)\n    \n    preds[:, 120+12:120+28] = (-test_q0002[:, 12:28]/1200)\n    pickle.dump(preds, open(f'preds_{epoch}.p', 'bw'))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = get_model()\nmodel.load_weights(f'/kaggle/input/leap-model-weights-lr/0p0-9epochs-E90.weights.h5')\nget_preds(model, 90)","metadata":{},"execution_count":null,"outputs":[]}]}