{"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":"none","dataSources":[{"sourceId":56537,"databundleVersionId":8877088,"sourceType":"competition"}],"dockerImageVersionId":30746,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# !pip install pathlib polars pandas matplotlib torch pytorch-lightning mlflow","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\ntorch.manual_seed(seed)\nfrom tqdm import tqdm\nfrom pathlib import Path\nimport polars as pl\nimport numpy as np\nimport pandas as pd\nfrom torch import nn\nfrom torch import optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader, random_split\nfrom pytorch_lightning import LightningDataModule, LightningModule, Trainer\nfrom pytorch_lightning.callbacks import ModelCheckpoint, LearningRateMonitor, EarlyStopping\nimport torch.nn.init as init\nimport os\nDEVICE=\"cuda\"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"k = 9\nbatch = 8\nbatch_slice = 128\nlr_rate = 4e-4; gamma_num = 0.9\nloss_num = 5\nHuberLoss_delta = 1\nmask = True\norder_extra_first = True\nseed = 42\n\nfinetune = False\ncheckpoint_path = \"/.pth\"\nif finetune:\n    ALL_FILES=[f'./input_low/{i}' for i in os.listdir('./input_low/')]\nelse:\n    ALL_FILES=[f'./input_aqua/{i}' for i in os.listdir('./input_aqua/')] + [f'./input_low/{i}' for i in os.listdir('./input_low/')]\nNUM_ROWS = len(ALL_FILES)\nDATA_MAP = {k:v for k,v in zip(range(NUM_ROWS),ALL_FILES)}","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"values_loaded = torch.load(\"./mean_std_values_col9_log_x.pth\")\nX_MEAN = values_loaded['X_MEAN']\nX_STD = values_loaded['X_STD']\nY_MEAN = values_loaded['Y_MEAN']\nY_STD = values_loaded['Y_STD']\n\nREPLACE_FROM = [f'ptend_q0002_{i}' for i in range(27)]\nREPLACE_TO = [f'state_q0002_{i}' for i in range(27)]\nFEATURE_NAMES = [f'state_t_{i}' for i in range(60)] + [f'state_q0001_{i}' for i in range(60)] + [f'state_q0002_{i}' for i in range(60)] + [f'state_q0003_{i}' for i in range(60)] + [f'state_u_{i}' for i in range(60)] + [f'state_v_{i}' for i in range(60)] + ['state_ps', 'pbuf_SOLIN', 'pbuf_LHFLX', 'pbuf_SHFLX', 'pbuf_TAUX', 'pbuf_TAUY', 'pbuf_COSZRS', 'cam_in_ALDIF', 'cam_in_ALDIR', 'cam_in_ASDIF', 'cam_in_ASDIR', 'cam_in_LWUP', 'cam_in_ICEFRAC', 'cam_in_LANDFRAC', 'cam_in_OCNFRAC', 'cam_in_SNOWHLAND'] + [f'pbuf_ozone_{i}' for i in range(60)] + [f'pbuf_CH4\nTARGET_NAMES = [f'ptend_t_{i}' for i in range(60)] + [f'ptend_q0001_{i}' for i in range(60)] + [f'ptend_q0002_{i}' for i in range(60)] + [f'ptend_q0003_{i}' for i in range(60)] + [f'ptend_u_{i}' for i in range(60)] + [f'ptend_v_{i}' for i in range(60)] + ['cam_out_NETSW', 'cam_out_FLWDS', 'cam_out_PRECSC', 'cam_out_PRECC', 'cam_out_SOLS', 'cam_out_SOLL', 'cam_out_SOLSD', 'cam_out_SOLLD']\n\nSTATE_T_IDX = list(range(0, 60))\nSTATE_T_IDX_2 = list(range(1, 61))\nSTATE_Q0001_IDX = list(range(60, 120))\nSTATE_Q0001_IDX_2 = list(range(61, 121))\nSTATE_U_IDX = list(range(240, 300))\nSTATE_U_IDX_2 = list(range(241, 301))\nSTATE_V_IDX = list(range(300, 360))\nSTATE_V_IDX_2 = list(range(301, 361))\nSTATE_Q0002_IDX = list(range(120, 180))\nSTATE_Q0003_IDX = list(range(180, 240))\nPBUF_OZONE_IDX = list(range(376, 436))\nPBUF_CH4_IDX = list(range(436, 496))\nPBUF_N2O_IDX = list(range(496, 556))\n\nTARGET_WEIGHTS = [30981.265271661872, 22502.432413914863, 18894.14713004499, 14514.244730542465, 10944.348069459196, 9065.01072024503, 9663.669038687454, 12688.557362943708, 19890.17226527665, 25831.37317235381, 33890.367561807274, 44122.94111025334, 59811.25595068309, 79434.07500078829, 107358.80916894016, 135720.8418348218, 149399.8411114814, 128492.95185325432, 91746.23687305572, 72748.76911097553, 66531.53596840335, 62932.30598423903, 56610.26874314136, 49473.14369220607, 43029.18495420936, 36912.67491908133, 31486.93117928144, 26898.072997215502, 23316.638282978325, 20459.73133196152, 18385.68309639014, 17111.405107656312, 16337.80991958771, 15857.759882318944, 15580.902485189716, 15497.59045982052, 15612.2556996736, 15797.88455410361, 15974.218740897895, 16130.395527176632, 16261.310866446129, 16371.892401608216, 16397.019695140876, 16325.463899570548, 16228.641108112768, 16191.809643436269, 16341.207925934068, 16645.711351490587, 17005.493716683693, 17430.29874509864, 17907.24023203076, 18431.55334008694, 19032.471309392287, 19701.355113141435, 20408.236605392685, 20967.20795006453, 21194.427318009974, 21088.521528526755, 19437.91555757985, 13677.902713248171, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 871528441401.8333, 1083221770553.0684, 147034752676.7702, 35556045575.13566, 35153369257.41337, 46086368691.51654, 24689305171.692936, 11343276593.440475, 5396624651.94418, 2449353007.641508, 1132225885.703891, 579547849.1340877, 330219246.7861086, 207613930.3131764, 144580292.27473342, 109933282.92266414, 88706603.092171, 73819777.54163922, 63615988.74519494, 57250262.292053565, 52976073.06761927, 49653169.17819005, 46544975.11484598, 43167606.9599748, 39724375.20499403, 36317177.25886468, 33057511.80930482, 29869089.497658804, 26982386.85583376, 24416235.17215712, 22273651.697369896, 20553426.04804544, 19216240.03357431, 18167694.44812838, 17501855.536957663, 17169938.630597908, 17005382.258644175, 16998475.26752617, 17082890.987979066, 17227982.77516062, 17445823.21630204, 17757404.421785507, 18346092.75160569, 19400573.66632694, 20506722.48296608, 22469648.380506545, 23432031.455169585, 26204163.40545158, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 1000000000000000.0, 3673829810926.31, 371405570725.2526, 14219163611.984406, 3001863018.1934915, 1432766589.9326108, 884599805.0283787, 560127980.1033351, 386052567.7087711, 287331851.051439, 222703657.59538063, 181069239.6264349, 154620864.3164144, 138093777.60284117, 126605828.89875436, 117967840.02553518, 111005814.39518328, 105186901.20678852, 100168133.0295481, 95568646.67416307, 91457433.39515457, 88871610.45308323, 88829796.26374224, 91398113.73291488, 96585131.67000748, 104507692.01463065, 115895119.998433, 131939701.08213414, 154492946.00677127, 183147918.17086875, 215151374.22324687, 247158314.6345976, 266792879.42215955, 279115128.29108113, 370541510.87006927, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 877670509694.7871, 1174826943136.8308, 1270605570069.038, 21727315470.5208, 3159456646.5437946, 1090653401.282219, 727967089.8459107, 384399548.9506704, 290787296.9451616, 232703218.45048887, 197467462.7577736, 174310890.8025987, 160536437.73297343, 153567098.77048483, 152120124.9453068, 153115566.6756177, 153955545.42558223, 153734675.21565756, 154798666.36905554, 163346213.58113608, 180013139.3707387, 200324358.8534948, 220754613.1646765, 241290935.478592, 262868932.2066308, 284448910.01847774, 305681084.4142859, 327605088.8575117, 350473296.7263526, 373964594.1196182, 398396925.8173239, 423528355.65716046, 450447055.544388, 478857006.4973163, 508200335.7126168, 537309657.5789208, 566854568.2904652, 594618842.9455439, 619715928.2391286, 641395460.8414665, 663290039.7810476, 689274894.631561, 718208866.3397261, 743951200.8024124, 761776104.2945968, 772911224.3082078, 804001144.8046833, 772448774.7758856, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 4613823.568205323, 1999308.9343799097, 904636.2296014762, 433823.6123842511, 207201.39055371704, 107836.09164720173, 57647.915219220784, 40606.52305039815, 47739.86647922776, 51669.35493930698, 56438.19768395407, 60447.45665200092, 65251.4153955275, 71920.88588011517, 78529.58115204438, 83422.30217897324, 87036.98552475807, 90389.72631774022, 93982.39165674087, 97578.0099352472, 101428.21366062944, 104630.69200130588, 105685.04322626138, 103962.58423268417, 99650.31670632094, 94290.49986206587, 89514.90144353417, 85905.45713126978, 82784.9857650212, 79152.28707014346, 74847.81017353121, 70378.81859610273, 65420.04643792357, 59953.75184604176, 54764.28281143022, 50362.51288353384, 46212.571031725325, 41997.52779088816, 37692.05148110484, 33834.73460995647, 31846.09764364542, 31934.145655397457, 31454.81247448105, 30105.4073072481, 26957.830283611693, 27760.04479210889, 29853.374336459365, 19133.428743715107, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 7619940.584531054, 3148394.472742347, 1308415.0022178134, 540515.7720745018, 215237.1053603881, 102546.7276372816, 68453.67122640925, 50692.59053608593, 51487.52043139844, 52104.76838400132, 54019.39151917722, 55856.02168787862, 60347.30240270209, 68990.96019017675, 79096.88768563846, 87574.33453690328, 94158.56052476274, 101903.63670531697, 111746.9753834774, 122460.65399236557, 132086.69387474353, 141041.48571028374, 146354.09441287292, 145953.09590059065, 139496.8007888401, 128508.85108217449, 116665.51769667884, 107458.39706309135, 100259.97236694951, 94108.98505029618, 88439.89456238014, 82734.9027659809, 77061.08621371102, 71333.5319243128, 65999.72532130677, 61798.9972058361, 58237.356419617165, 54715.10266341248, 50825.84431702935, 46059.17688689915, 40740.26050401376, 36335.80228304863, 33981.57568605091, 33589.7143390849, 33988.88524112733, 36272.9364507092, 41183.34413717943, 29194.12369278645, 0.0040536134869726, 0.0138824238058072, 135129884.5084534, 12219717.5342461, 0.0090705273332672, 0.0085898851680217, 0.0215368188774867, 0.0336321308942602]\nNEW_TARGET_WEIGHTS = [1] * 60 + [0] * 12 + [1] * 48 + [0] * 12 + [1] * 48 + [0] * 12 + [1] * 48 + [0] * 12 + [1] * 48 + [0] * 12 + [1] * 56","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_x(x):\n    temp_diff1 = x[:,  STATE_T_IDX] - (x[:,STATE_T_IDX_2])\n    temp_diff2 = x[:,  STATE_Q0001_IDX] - (x[:,STATE_Q0001_IDX_2])\n    wind_diff1 = x[:,  STATE_U_IDX] - (x[:,STATE_U_IDX_2])\n    wind_diff2 = x[:,  STATE_V_IDX] - (x[:,STATE_V_IDX_2])\n    temp_humid = x[:, STATE_T_IDX] / (x[:,STATE_Q0001_IDX])\n    air_total = x[:, PBUF_OZONE_IDX] + (x[:,PBUF_CH4_IDX]) + x[:, PBUF_N2O_IDX]\n    moisture = ((x[:, STATE_Q0001_IDX])**2) * ((x[:,STATE_U_IDX])**2 + (x[:,STATE_V_IDX])**2)\n    liq_partition = x[:, STATE_Q0002_IDX] / (x[:,STATE_Q0002_IDX] + x[:,STATE_Q0003_IDX])\n    imbalance = (x[:,STATE_Q0002_IDX] - x[:,STATE_Q0003_IDX]) / (x[:,STATE_Q0002_IDX] + x[:,STATE_Q0003_IDX])\n    additional_features = torch.nan_to_num(torch.cat([temp_diff1, temp_diff2, wind_diff1, wind_diff2, temp_humid, air_total, moisture, liq_partition, imbalance], -1),0)\n    return torch.cat([x, additional_features], -1)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LeapDataset(Dataset):\n    def __init__(self, x_features, y_features, y_weights):\n        self.x_features =x_features\n        self.x_split = len(x_features)\n        self.y_features = y_features\n        self.y_weights = y_weights\n\n    def __getitem__(self, idx):\n        data = torch.load(DATA_MAP[idx])\n        x, y = torch.split(data, self.x_split, dim=1)\n        y = y * self.y_weights\n        x = preprocess_x(x)\n        x_mean = X_MEAN; x_std = X_STD; y_mean = Y_MEAN; y_std = Y_STD \n        x = (x - x_mean) / x_std; y = (y - y_mean) / y_std\n        x = x.to(torch.float32)\n\n        if order_extra_first:\n            x = torch.cat([torch.cat(x[: , 556:].reshape(batch_slice, k, 60), [x[: , :360].reshape(batch_slice, 6, 60), x[: , 376:556].reshape(batch_slice, 3, 60)], dim=1)\n                .permute(0,2, 1), x[: , 360:376].unsqueeze(1).repeat(1,60, 1).view(batch_slice, 60, 16),],-1,)\n        else:\n            x = torch.cat([torch.cat([x[: , :360].reshape(batch_slice, 6, 60), x[: , 556:].reshape(batch_slice, k, 60), x[: , 376:556].reshape(batch_slice, 3, 60)], dim=1)\n                .permute(0,2, 1), x[: , 360:376].unsqueeze(1).repeat(1,60, 1).view(batch_slice, 60, 16),],-1,)\n                \n        y = y.to(torch.float32)\n        return x, y\n        \n    def __len__(self):\n        return NUM_ROWS\n\nds_data = LeapDataset(x_features=FEATURE_NAMES, y_features=TARGET_NAMES, y_weights=torch.tensor(TARGET_WEIGHTS))\nds_train, ds_valid = random_split(ds_data, [len(ds_data)-2000, 2000])\ntrain_loader = DataLoader(ds_train, batch_size=batch, shuffle=True, drop_last=True, pin_memory=True, persistent_workers=True, num_workers=32, prefetch_factor=4)\nvalid_loader = DataLoader(ds_valid, batch_size=1, shuffle=False, drop_last=False,pin_memory=False, persistent_workers=True,num_workers=16)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def r2_score(y_pred, y_true):\n    ss_res = torch.sum((y_true - y_pred) ** 2)\n    ss_tot = torch.sum((y_true - torch.mean(y_true)) ** 2)\n    r2 = 1 - ss_res / ss_tot\n    return r2.item()\n\nWEIGHT_MASK = torch.tensor(NEW_TARGET_WEIGHTS)\nADJUSTMENT_COLUMNS = [f\"ptend_q0002_{i}\" for i in range(28)]\nADJUSTMENT_MASK=[]\n\nfor col, wt in zip(TARGET_NAMES, NEW_TARGET_WEIGHTS):\n    if wt==0 and col in ADJUSTMENT_COLUMNS:\n        ADJUSTMENT_MASK.append(0)\n    else:\n        ADJUSTMENT_MASK.append(1)\n\nADJUSTMENT_MASK=torch.tensor(ADJUSTMENT_MASK)\nMASK = torch.tensor(NEW_TARGET_WEIGHTS)\nJOINT_MASK = MASK * ADJUSTMENT_MASK","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MLP(LightningModule):\n    def __init__(self, dims):\n        super().__init__()\n        layers = []\n        for i in range(len(dims) - 2):\n            layers.append(nn.Linear(dims[i], dims[i + 1]))\n            layers.append(nn.ReLU())\n        layers.append(nn.Linear(dims[-2], dims[-1]))\n        self.network = nn.Sequential(*layers)\n        self._initialize_weights()\n\n    def _initialize_weights(self):\n        for m in self.network:\n            if isinstance(m, nn.Linear):\n                init.xavier_uniform_(m.weight)\n                if m.bias is not None:\n                    init.zeros_(m.bias)\n            elif isinstance(m, nn.BatchNorm1d):\n                init.ones_(m.weight)\n                init.zeros_(m.bias)\n\n    def forward(self, x):\n        return self.network(x)\n\nclass LeapModel(LightningModule):\n    def __init__(self, input_dim, hidden_dim, output_dim, num_layers=2):\n        super().__init__()\n        self.encoder = MLP([input_dim * 60, input_dim * 15, input_dim * 8, input_dim * 4])\n        self.decoder = MLP([input_dim * 4, input_dim * 8, input_dim * 15, input_dim * 60])\n        self.lstm = nn.LSTM(input_dim * 2, hidden_dim, num_layers, batch_first=True, dropout=0.07, bidirectional=True)\n        self.lstm2 = nn.GRU(hidden_dim * 2, 16, batch_first=True, dropout=0.03, bidirectional=True)\n        self.fc_lstm = nn.Linear((32) * 60 , output_dim)\n        self.criterion = nn.HuberLoss(delta=HuberLoss_delta)\n\n    def forward(self, x):\n        x_0 = self.encoder(x.view(x.size(0), -1))\n        x_1 = self.decoder(x_0)\n        x_1 = x_1.view(x.size(0), 60, -1)\n        lstm_out, _ = self.lstm(torch.cat([x, x_1], dim = -1))\n        lstm_out, _ = self.lstm2(lstm_out)\n        lstm_out = lstm_out.contiguous().view(x.size(0), -1) \n        output = self.fc_lstm(lstm_out)\n        return output\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        y = y.view(y.size(1)*y.size(0), y.size(2))\n        x = x.view(x.size(1)*x.size(0), x.size(2), x.size(3))\n        y_pred = self(x)\n\n        if mask:\n            loss = self.criterion(y_pred[:, JOINT_MASK==1], y[:, JOINT_MASK==1])\n        else:\n            loss = self.criterion(y_pred, y)\n\n        loss *= loss_multiple_num\n        self.log('train_loss', loss, prog_bar=True, on_epoch=True, on_step=False, logger=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        x, y = batch\n        y = y.view(y.size(1)*y.size(0), y.size(2))\n        x = x.view(x.size(1)*x.size(0), x.size(2), x.size(3))\n        y_pred = self(x)\n        y_std = Y_STD.to(y.device)\n        y_mean = Y_MEAN.to(y.device)\n        y = (y * y_std) + y_mean\n        y_pred[:, y_std < (1.1 * 1e-6)] = 0\n        y_pred = (y_pred * y_std) + y_mean\n        val_score = r2_score(y_pred, y)\n        self.log('val_score', val_score, prog_bar=True, on_epoch=True, on_step=False, logger=True)\n        \n        if mask:\n            y_pred[:, ADJUSTMENT_MASK==0] = y[:, ADJUSTMENT_MASK==0]\n            y_pred[:, MASK==0] = 0\n            y[:, MASK==0] = 0\n            masked_val_score = r2_score(y_pred, y)\n            self.log('masked_val_score', masked_val_score, prog_bar=True, on_epoch=True, on_step=False, logger=True)\n        return val_score\n\n    def configure_optimizers(self):\n        optimizer = optim.AdamW(self.parameters(), lr=lr_rate, weight_decay=5e-4)\n        milestones = [1, 2, 3, 5, 6, 7, 8, 9, 11, 13, 15, 17, 19] \n        gamma = gamma_num\n        scheduler = optim.lr_scheduler.MultiStepLR(optimizer, milestones=milestones, gamma=gamma)\n        return [optimizer], [scheduler]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = LeapModel(input_dim=25+k, hidden_dim=512, output_dim=368, num_layers=3)\nevery_epoch_checkpoint_callback = ModelCheckpoint(every_n_epochs=1, filename=\"{epoch:02d}-{val_score:.4f}\", dirpath=\"/output\", save_top_k=-1)\n\ntrainer = Trainer(\n    max_epochs=20,\n    accelerator=\"gpu\",\n    devices=1,\n    enable_checkpointing=True,\n    precision=16,\n    callbacks=[every_epoch_checkpoint_callback],\n    val_check_interval=0.5)\n\nif finetune:\n    checkpoint = torch.load(checkpoint_path, map_location=DEVICE)\n    model.load_state_dict(checkpoint['state_dict'])\n    \ntorch.set_float32_matmul_precision('medium')\ntrainer.fit(model, train_loader, valid_loader)","metadata":{},"execution_count":null,"outputs":[]}]}