{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":11553390,"sourceType":"competition"},{"sourceId":10933126,"sourceType":"datasetVersion","datasetId":6798383},{"sourceId":230671906,"sourceType":"kernelVersion"},{"sourceId":230731739,"sourceType":"kernelVersion"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Setup","metadata":{}},{"cell_type":"code","source":"!pip install -q lightning bitsandbytes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:23.684438Z","iopub.execute_input":"2025-03-31T13:10:23.684828Z","iopub.status.idle":"2025-03-31T13:10:27.191525Z","shell.execute_reply.started":"2025-03-31T13:10:23.684766Z","shell.execute_reply":"2025-03-31T13:10:27.190509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 42\nimport os\nimport numpy as np\nfrom numpy import random as np_rnd\nimport random as rnd\nimport pandas as pd\nimport pickle\nimport gc\nimport time\n\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence\nfrom PIL import Image\nfrom bitsandbytes.optim import AdamW8bit\nfrom transformers import get_polynomial_decay_schedule_with_warmup\nimport lightning as L\nfrom lightning.pytorch.callbacks import ModelCheckpoint, LearningRateMonitor\nfrom lightning.pytorch import loggers as pl_loggers\nfrom lightning.pytorch.callbacks.early_stopping import EarlyStopping\nfrom conformer import Conformer\nimport sklearn.metrics as skl_metrics\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ndevice","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:27.193033Z","iopub.execute_input":"2025-03-31T13:10:27.193333Z","iopub.status.idle":"2025-03-31T13:10:27.202768Z","shell.execute_reply.started":"2025-03-31T13:10:27.193297Z","shell.execute_reply":"2025-03-31T13:10:27.202158Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed=42):\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    # python random\n    rnd.seed(seed)\n    # numpy random\n    np_rnd.seed(seed)\n    # RAPIDS random\n    try:\n        cupy.random.seed(seed)\n    except:\n        pass\n    # tf random\n    try:\n        tf_rnd.set_seed(seed)\n    except:\n        pass\n    # pytorch random\n    try:\n        torch.backends.cudnn.benchmark = False\n        torch.backends.cudnn.deterministic = True\n        torch.manual_seed(seed)\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n    except:\n        pass\n\ndef pickleIO(obj, src, op=\"r\"):\n    if op==\"w\":\n        with open(src, op + \"b\") as f:\n            pickle.dump(obj, f)\n    elif op==\"r\":\n        with open(src, op + \"b\") as f:\n            tmp = pickle.load(f)\n        return tmp\n    else:\n        print(\"unknown operation\")\n        return obj\n    \ndef createFolder(directory):\n    try:\n        if not os.path.exists(directory):\n            os.makedirs(directory)\n    except OSError:\n        print('Error: Creating directory. ' + directory)\n\ndef findIdx(data_x, col_names):\n    return [int(i) for i, j in enumerate(data_x) if j in col_names]\n\ndef diff(first, second):\n    second = set(second)\n    return [item for item in first if item not in second]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:27.204481Z","iopub.execute_input":"2025-03-31T13:10:27.204689Z","iopub.status.idle":"2025-03-31T13:10:27.221570Z","shell.execute_reply.started":"2025-03-31T13:10:27.204670Z","shell.execute_reply":"2025-03-31T13:10:27.220794Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    debug = False\n    dp_version = \"a1\"\n    n_folds = 5\n    max_seq = 1024\n    epochs = 2 if debug else 20\n    early_stopping_rounds = 5\n    batch_size = 4 if debug else 16\n    eta = 2e-4\n    weight_decay = 1e-2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:27.223077Z","iopub.execute_input":"2025-03-31T13:10:27.223341Z","iopub.status.idle":"2025-03-31T13:10:27.240537Z","shell.execute_reply.started":"2025-03-31T13:10:27.223318Z","shell.execute_reply":"2025-03-31T13:10:27.239742Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seed_everything(SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:27.241275Z","iopub.execute_input":"2025-03-31T13:10:27.241523Z","iopub.status.idle":"2025-03-31T13:10:27.257563Z","shell.execute_reply.started":"2025-03-31T13:10:27.241502Z","shell.execute_reply":"2025-03-31T13:10:27.256796Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loading data","metadata":{}},{"cell_type":"code","source":"df_full = pickleIO(None, f\"/kaggle/input/srna3d-dp-{CFG.dp_version}/df_full.pkl\", \"r\")\nlabel_full = pickleIO(None, f\"/kaggle/input/srna3d-dp-{CFG.dp_version}/label_full.pkl\", \"r\")\ndf_test = pickleIO(None, f\"/kaggle/input/srna3d-dp-{CFG.dp_version}/df_test.pkl\", \"r\")\nlabel_test = pickleIO(None, f\"/kaggle/input/srna3d-dp-{CFG.dp_version}/label_test.pkl\", \"r\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:27.258604Z","iopub.execute_input":"2025-03-31T13:10:27.258930Z","iopub.status.idle":"2025-03-31T13:10:27.365268Z","shell.execute_reply.started":"2025-03-31T13:10:27.258898Z","shell.execute_reply":"2025-03-31T13:10:27.364551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_full","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:27.366181Z","iopub.execute_input":"2025-03-31T13:10:27.366390Z","iopub.status.idle":"2025-03-31T13:10:27.382064Z","shell.execute_reply.started":"2025-03-31T13:10:27.366370Z","shell.execute_reply":"2025-03-31T13:10:27.381136Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_full","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:27.383067Z","iopub.execute_input":"2025-03-31T13:10:27.383388Z","iopub.status.idle":"2025-03-31T13:10:27.406100Z","shell.execute_reply.started":"2025-03-31T13:10:27.383354Z","shell.execute_reply":"2025-03-31T13:10:27.405298Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"residues = pickleIO(None, f\"/kaggle/input/srna3d-dp-{CFG.dp_version}/residues.pkl\", \"r\")\nresidues","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:27.408001Z","iopub.execute_input":"2025-03-31T13:10:27.408281Z","iopub.status.idle":"2025-03-31T13:10:27.428500Z","shell.execute_reply.started":"2025-03-31T13:10:27.408259Z","shell.execute_reply":"2025-03-31T13:10:27.427738Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define helper functions","metadata":{}},{"cell_type":"code","source":"class SequenceConFormer(L.LightningModule):\n    def __init__(self, df_coord, conformer_params, vocab_size, num_classes, eta, weight_decay, num_training_steps):\n        super().__init__()\n        self.df_coord = df_coord\n        self.embedding = nn.Embedding(vocab_size, conformer_params[\"dim\"])\n        self.conformer = Conformer(**conformer_params)\n        self.regressor = nn.Linear(conformer_params[\"dim\"], num_classes)\n        self.eta = eta\n        self.weight_decay = weight_decay\n        self.num_training_steps = num_training_steps\n        self.criterion = nn.MSELoss()\n\n    def forward(self, x):\n        x = self.embedding(x)\n        x = self.conformer(x)\n        x = self.regressor(x)\n        return x\n\n    def calc_loss(self, outputs, labels, seq_locs, df_coord):\n        losses = []\n        for output, label, seq_loc in zip (outputs, labels, seq_locs):\n            targets = torch.tensor(df_coord.loc[label, [\"x_1\", \"y_1\", \"z_1\"]].values)[seq_loc[0]:seq_loc[1]].to(device)\n            losses.append(self.criterion(output[:len(targets)], targets))\n        loss = sum(losses) / len(losses)\n        return loss\n        \n    def training_step(self, batch):\n        outputs = self(batch[\"input_ids\"])\n        loss = self.calc_loss(outputs, batch[\"label\"], batch[\"seq_loc\"], self.df_coord)\n        self.log(\"train_loss\", loss)\n        return loss\n\n    def validation_step(self, batch):\n        outputs = self(batch[\"input_ids\"])\n        loss = self.calc_loss(outputs, batch[\"label\"], batch[\"seq_loc\"], self.df_coord)\n        self.log(\"val_loss\", loss)\n        return loss\n\n    def get_optimizer_params(self, eta, weight_decay):\n        no_decay = [\"bias\", \"LayerNorm.bias\", \"LayerNorm.weight\"]\n        optimizer_parameters = [\n            # apply weight decay\n            {'params': [p for n, p in self.named_parameters() if not any(nd in n for nd in no_decay)],\n            'lr': eta, 'weight_decay': weight_decay},\n            # don't apply weight decay for LayerNormalization layer\n            {'params': [p for n, p in self.named_parameters() if any(nd in n for nd in no_decay)],\n            'lr': eta, 'weight_decay': 0.0},\n        ]\n        return optimizer_parameters\n\n    def get_scheduler(self, optimizer, num_warmup_steps, num_training_steps):\n        scheduler = get_polynomial_decay_schedule_with_warmup(\n            optimizer, num_warmup_steps=num_warmup_steps, num_training_steps=num_training_steps, power=1.0, lr_end=1e-7\n        )\n        return scheduler\n\n    def configure_optimizers(self):\n        optimizer_parameters = self.get_optimizer_params(\n            eta=self.eta,\n            weight_decay=self.weight_decay\n        )   \n        optimizer = AdamW8bit(optimizer_parameters, lr=self.eta, weight_decay=self.weight_decay)\n        scheduler = self.get_scheduler(\n            optimizer,\n            num_warmup_steps=0,\n            num_training_steps=self.num_training_steps\n        )\n        return {\"optimizer\": optimizer, \"lr_scheduler\": {\"scheduler\": scheduler, \"interval\": \"step\"}}\n\n@torch.no_grad()\ndef inference(model, dl):\n    model.eval()\n    output = []\n    for batch in dl:\n        output.append(model(batch[\"input_ids\"].to(device)).squeeze(0).detach().cpu().numpy())\n    return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:27.429616Z","iopub.execute_input":"2025-03-31T13:10:27.429945Z","iopub.status.idle":"2025-03-31T13:10:27.450695Z","shell.execute_reply.started":"2025-03-31T13:10:27.429913Z","shell.execute_reply":"2025-03-31T13:10:27.450036Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, features, labels, seq_locs):\n        self.features = features\n        self.labels = labels\n        self.seq_locs = seq_locs\n\n    def __len__(self):\n        return len(self.features)\n\n    def __getitem__(self, idx):\n        return {\"input_ids\": self.features[idx], \"label\": self.labels[idx], \"seq_loc\": self.seq_locs[idx]}\n\ndef collate_fn(samples):\n    batch = {\n        \"input_ids\": pad_sequence([torch.tensor(sample[\"input_ids\"]) for sample in samples], batch_first=True),\n        \"label\": [sample[\"label\"] for sample in samples],\n        \"seq_loc\": [sample[\"seq_loc\"] for sample in samples],\n    }\n    return batch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:27.451496Z","iopub.execute_input":"2025-03-31T13:10:27.451798Z","iopub.status.idle":"2025-03-31T13:10:27.468758Z","shell.execute_reply.started":"2025-03-31T13:10:27.451737Z","shell.execute_reply":"2025-03-31T13:10:27.468012Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"architecture_path = \"./\"\nfixed_params = {\n    \"conformer_params\": {\n        \"dim\": 512,\n        \"depth\": 6,\n        \"dim_head\": 64,\n        \"heads\": 8,\n        \"ff_mult\": 4,\n        \"conv_expansion_factor\": 2,\n        \"conv_kernel_size\": 31,\n        \"attn_dropout\": 0.1,\n        \"ff_dropout\": 0.1,\n        \"conv_dropout\": 0.1\n    },\n    \"eta\": CFG.eta,\n    \"weight_decay\": CFG.weight_decay,\n}\ntest_x = df_test[\"sequence\"].apply(lambda x: pd.Series(list(x)).map(residues).fillna(0.0).astype(\"int64\").to_list())\ntest_y =  df_test[\"target_id\"]\ntest_seq_locs = df_test[\"seq_loc\"]\ntest_dl = DataLoader(CustomDataset(test_x.tolist(), test_y.tolist(), test_seq_locs.tolist()), batch_size=1, collate_fn=collate_fn, shuffle=False)\n\nfold_pred = []\nfold_score = []\nfor fold in range(CFG.n_folds):\n    print(f\"\\n=== FOLD {fold} ===\")\n    start_time = time.time()\n    # split train & valid\n    df_train = df_full.loc[df_full[\"fold_target\"] != 0].reset_index(drop=True)\n    train_x = df_train[\"sequence\"].apply(lambda x: pd.Series(list(x)).map(residues).fillna(0.0).astype(\"int64\").to_list())\n    train_y = df_train[\"target_id\"]\n    train_seq_locs = df_train[\"seq_loc\"]\n    df_valid = df_full.loc[df_full[\"fold_target\"] == 0].reset_index(drop=True)\n    valid_x = df_valid[\"sequence\"].apply(lambda x: pd.Series(list(x)).map(residues).fillna(0.0).astype(\"int64\").to_list())\n    valid_y =  df_valid[\"target_id\"]\n    valid_seq_locs = df_valid[\"seq_loc\"]\n    if CFG.debug:\n        train_x = train_x.iloc[:100]\n        train_y = train_y.iloc[:100]\n        valid_x = valid_x.iloc[:100]\n        valid_y = valid_y.iloc[:100]\n    print(\"shape info ->\", train_x.shape, train_y.shape, valid_x.shape, valid_y.shape)\n    # create dataloader\n    train_dl = DataLoader(CustomDataset(train_x.tolist(), train_y.tolist(), train_seq_locs.tolist()), batch_size=CFG.batch_size, collate_fn=collate_fn, shuffle=True, drop_last=True)\n    valid_dl = DataLoader(CustomDataset(valid_x.tolist(), valid_y.tolist(), valid_seq_locs.tolist()), batch_size=CFG.batch_size, collate_fn=collate_fn, shuffle=False)\n    # create model\n    model_params = fixed_params.copy()\n    model_params[\"vocab_size\"] = len(residues) + 1\n    model_params[\"num_classes\"] = 3\n    model_params[\"num_training_steps\"] = len(train_dl) * CFG.epochs\n    model = SequenceConFormer(**model_params, df_coord=label_full)\n    model.to(device)\n    # training\n    checkpoint_callback = ModelCheckpoint(\n        monitor='val_loss',\n        dirpath=os.path.join(architecture_path, f\"ckpt/fold{fold}\"),\n        save_last=True,\n    )\n    earlystopping_callback = EarlyStopping(\n        monitor='val_loss',\n        patience=CFG.early_stopping_rounds,\n        verbose=True,\n        mode='min'\n    )\n    lr_callback = LearningRateMonitor(logging_interval='epoch')\n    trainer = L.Trainer(\n        max_epochs=CFG.epochs, \n        accelerator=\"gpu\",\n        devices=[0],\n        precision=16,\n        logger=pl_loggers.CSVLogger(save_dir=os.path.join(architecture_path, f\"logs/fold{fold}\")),\n        callbacks=[checkpoint_callback, lr_callback, earlystopping_callback],\n    )\n    trainer.fit(model, train_dl, valid_dl)\n    model.to(device)\n    model.load_state_dict(torch.load(checkpoint_callback.best_model_path, map_location=device)[\"state_dict\"])\n    # validation\n    y_pred = inference(model, test_dl)\n    # evaluation\n    fold_score.append({\n        \"loss\": model.calc_loss([torch.tensor(i).to(device) for i in y_pred], test_y, test_seq_locs, label_test).item(),\n    })\n    print(\"[SCORE]\")\n    print(pd.Series(fold_score[-1]))\n    # save data\n    pickleIO(\n        {\n            \"model_params\": model_params,\n        },\n        os.path.join(architecture_path, f\"fold{fold}_output.pkl\"),\n        \"w\",\n    )\n    del train_x, train_y, train_seq_locs, valid_x, valid_y, valid_seq_locs, train_dl, valid_dl, model, trainer\n    torch.cuda.empty_cache()    \n    gc.collect()\n    print(\"Elapsed time: {:.2f}s\".format(time.time() - start_time))\n    \npickleIO(fold_pred, os.path.join(architecture_path, \"fold_pred.pkl\"), \"w\")\npickleIO(fold_score, os.path.join(architecture_path, \"fold_score.pkl\"), \"w\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T13:10:27.469600Z","iopub.execute_input":"2025-03-31T13:10:27.469850Z","execution_failed":"2025-03-31T13:11:17.329Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## CV score","metadata":{}},{"cell_type":"code","source":"df_score = pd.DataFrame(fold_score)\ndf_score.loc[\"average\"] = df_score.mean()\ndf_score.round(5)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-31T13:11:17.329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}