{"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":"none","dataSources":[{"sourceId":87793,"databundleVersionId":11553390,"sourceType":"competition"},{"sourceId":10933126,"sourceType":"datasetVersion","datasetId":6798383},{"sourceId":230671906,"sourceType":"kernelVersion"},{"sourceId":230748062,"sourceType":"kernelVersion"},{"sourceId":230888557,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Setup","metadata":{}},{"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 glob\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 transformers import get_polynomial_decay_schedule_with_warmup\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-04-01T00:44:43.600109Z","iopub.execute_input":"2025-04-01T00:44:43.600520Z","iopub.status.idle":"2025-04-01T00:44:48.119931Z","shell.execute_reply.started":"2025-04-01T00:44:43.600484Z","shell.execute_reply":"2025-04-01T00:44:48.118755Z"}},"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-04-01T00:44:48.121245Z","iopub.execute_input":"2025-04-01T00:44:48.121800Z","iopub.status.idle":"2025-04-01T00:44:48.131566Z","shell.execute_reply.started":"2025-04-01T00:44:48.121768Z","shell.execute_reply":"2025-04-01T00:44:48.130369Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    debug = False\n    dp_version = \"a1\"\n    architecture_version = \"b1\"\n    n_folds = 5\n    max_seq = 1024","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T00:44:48.133608Z","iopub.execute_input":"2025-04-01T00:44:48.133989Z","iopub.status.idle":"2025-04-01T00:44:48.157006Z","shell.execute_reply.started":"2025-04-01T00:44:48.133944Z","shell.execute_reply":"2025-04-01T00:44:48.155851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seed_everything(SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T00:44:48.158588Z","iopub.execute_input":"2025-04-01T00:44:48.159004Z","iopub.status.idle":"2025-04-01T00:44:48.180986Z","shell.execute_reply.started":"2025-04-01T00:44:48.158962Z","shell.execute_reply":"2025-04-01T00:44:48.179884Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loading data","metadata":{}},{"cell_type":"code","source":"df_test = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")\ndf_test.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T00:44:48.182224Z","iopub.execute_input":"2025-04-01T00:44:48.182524Z","iopub.status.idle":"2025-04-01T00:44:48.211773Z","shell.execute_reply.started":"2025-04-01T00:44:48.182498Z","shell.execute_reply":"2025-04-01T00:44:48.210589Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"target_scaler = pickleIO(None, f\"/kaggle/input/srna3d-dp-{CFG.dp_version}/target_scaler.pkl\", \"r\")\nresidues = pickleIO(None, f\"/kaggle/input/srna3d-dp-{CFG.dp_version}/residues.pkl\", \"r\")\nresidues","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T00:44:48.212907Z","iopub.execute_input":"2025-04-01T00:44:48.213343Z","iopub.status.idle":"2025-04-01T00:44:48.223370Z","shell.execute_reply.started":"2025-04-01T00:44:48.213303Z","shell.execute_reply":"2025-04-01T00:44:48.222005Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define helper functions","metadata":{}},{"cell_type":"code","source":"class SequenceConFormer(nn.Module):\n    def __init__(self, conformer_params, vocab_size, num_classes):\n        super().__init__()\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\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@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-04-01T00:44:48.224297Z","iopub.execute_input":"2025-04-01T00:44:48.224588Z","iopub.status.idle":"2025-04-01T00:44:48.236156Z","shell.execute_reply.started":"2025-04-01T00:44:48.224552Z","shell.execute_reply":"2025-04-01T00:44:48.235132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, features):\n        self.features = features\n\n    def __len__(self):\n        return len(self.features)\n\n    def __getitem__(self, idx):\n        return {\"input_ids\": self.features[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    }\n    return batch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T00:44:48.239296Z","iopub.execute_input":"2025-04-01T00:44:48.239654Z","iopub.status.idle":"2025-04-01T00:44:48.259953Z","shell.execute_reply.started":"2025-04-01T00:44:48.239625Z","shell.execute_reply":"2025-04-01T00:44:48.258979Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"test_x = df_test[\"sequence\"].apply(lambda x: pd.Series(list(x)).map(residues).fillna(0.0).astype(\"int64\").to_list())\ntest_dl = DataLoader(CustomDataset(test_x.tolist()), batch_size=1, collate_fn=collate_fn, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T00:44:48.261317Z","iopub.execute_input":"2025-04-01T00:44:48.261612Z","iopub.status.idle":"2025-04-01T00:44:48.287271Z","shell.execute_reply.started":"2025-04-01T00:44:48.261588Z","shell.execute_reply":"2025-04-01T00:44:48.286307Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"architecture_path = f\"/kaggle/input/srna3d-training-conformer-{CFG.dp_version}-{CFG.architecture_version}\"\n\ntest_pred = []\nfor fold in range(CFG.n_folds):\n    print(f\"\\n=== FOLD {fold} ===\")\n    start_time = time.time()\n    fold_output = pickleIO(None, os.path.join(architecture_path, f\"fold{fold}_output.pkl\"), \"r\")\n    # load weights\n    model_params = fold_output[\"model_params\"]\n    model = SequenceConFormer(\n        conformer_params=model_params[\"conformer_params\"],\n        vocab_size=model_params[\"vocab_size\"],\n        num_classes=model_params[\"num_classes\"],\n    )\n    ckpt_path = glob.glob(os.path.join(architecture_path, f\"ckpt/fold{fold}/epoch*.ckpt\"))[0]\n    model.to(device)\n    model.load_state_dict(torch.load(ckpt_path, map_location=device)[\"state_dict\"])\n    # inference\n    y_pred = inference(model, test_dl)\n    test_pred.append(y_pred)\n    del model\n    torch.cuda.empty_cache()\n    gc.collect()\n    print(\"Elapsed time: {:.2f}s\".format(time.time() - start_time))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T00:44:48.288501Z","iopub.execute_input":"2025-04-01T00:44:48.288807Z","iopub.status.idle":"2025-04-01T00:45:09.958445Z","shell.execute_reply.started":"2025-04-01T00:44:48.288778Z","shell.execute_reply":"2025-04-01T00:45:09.957333Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"submission = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/sample_submission.csv\")\nsubmission[\"target_id\"] = submission[\"ID\"].apply(lambda x: x.split(\"_\")[0])\nsubmission = submission.set_index(\"target_id\")\nsubmission.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T00:56:41.498523Z","iopub.execute_input":"2025-04-01T00:56:41.498925Z","iopub.status.idle":"2025-04-01T00:56:41.534236Z","shell.execute_reply.started":"2025-04-01T00:56:41.498893Z","shell.execute_reply":"2025-04-01T00:56:41.533037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for fold in range(len(test_pred)):\n    preds = test_pred[fold]\n    for target_id, pred in zip(df_test[\"target_id\"], preds):\n        pred = target_scaler.inverse_transform(pred)\n        submission.loc[target_id, [f\"x_{fold+1}\", f\"y_{fold+1}\", f\"z_{fold+1}\"]] = pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T00:56:41.670709Z","iopub.execute_input":"2025-04-01T00:56:41.671097Z","iopub.status.idle":"2025-04-01T00:56:41.788709Z","shell.execute_reply.started":"2025-04-01T00:56:41.671061Z","shell.execute_reply":"2025-04-01T00:56:41.787217Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T00:45:10.031440Z","iopub.execute_input":"2025-04-01T00:45:10.031751Z","iopub.status.idle":"2025-04-01T00:45:10.059515Z","shell.execute_reply.started":"2025-04-01T00:45:10.031726Z","shell.execute_reply":"2025-04-01T00:45:10.058337Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission.reset_index(drop=True).to_csv(\"submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T00:45:10.060558Z","iopub.execute_input":"2025-04-01T00:45:10.060947Z","iopub.status.idle":"2025-04-01T00:45:10.115072Z","shell.execute_reply.started":"2025-04-01T00:45:10.060906Z","shell.execute_reply":"2025-04-01T00:45:10.113807Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}