{"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":"gpu","dataSources":[{"sourceId":67356,"databundleVersionId":8006601,"sourceType":"competition"},{"sourceId":7856946,"sourceType":"datasetVersion","datasetId":4608382}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Update\n|version|update|cv|lb|\n|:--|:--|--:|--:|\n|v3|-|-|0.343|\n|v4|increase N_ROWS from 90M to 180M|0.573|0.465|\n|v5|increase N_SAMPLES from 1M to 2M<br>change to a larger model (5M -> 10M)|0.622|0.486|\n|v8|Add normalization commented by [@hengck23](https://www.kaggle.com/hengck23)|0.629|0.496|","metadata":{}},{"cell_type":"code","source":"!pip install rdkit\n!pip install -U /kaggle/input/lightning-2-2-1/lightning-2.2.1-py3-none-any.whl","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-23T04:24:15.846197Z","iopub.execute_input":"2024-04-23T04:24:15.846535Z","iopub.status.idle":"2024-04-23T04:24:45.658021Z","shell.execute_reply.started":"2024-04-23T04:24:15.846507Z","shell.execute_reply":"2024-04-23T04:24:45.657007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pathlib import Path\n\nimport numpy as np\nimport polars as pl\nfrom sklearn.model_selection import train_test_split\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchmetrics import AveragePrecision\nimport lightning as L\nfrom lightning.pytorch.callbacks import (\n    EarlyStopping,\n    ModelCheckpoint,\n    TQDMProgressBar,\n)\nfrom transformers import AutoConfig, AutoTokenizer, AutoModel, DataCollatorWithPadding\nimport datasets\nfrom rdkit import Chem","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:24:45.660443Z","iopub.execute_input":"2024-04-23T04:24:45.660942Z","iopub.status.idle":"2024-04-23T04:25:04.515527Z","shell.execute_reply.started":"2024-04-23T04:24:45.660900Z","shell.execute_reply":"2024-04-23T04:25:04.514518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Hyper-Parameters","metadata":{}},{"cell_type":"code","source":"DEBUG = False\nNORMALIZE = True\nN_ROWS = 180_000_000\nassert N_ROWS is None or N_ROWS % 3 == 0\nif DEBUG:\n    N_SAMPLES = 10_000\nelse:\n    N_SAMPLES = 2_000_000\nPROTEIN_NAMES = [\"BRD4\", \"HSA\", \"sEH\"]\ndata_dir = Path(\"/kaggle/input/leash-BELKA\")\nmodel_name = \"DeepChem/ChemBERTa-10M-MTR\"\nbatch_size = 256\ntrainer_params = {\n  \"max_epochs\": 5,\n  \"enable_progress_bar\": True,\n  \"accelerator\": \"auto\",\n  \"precision\": \"16-mixed\",\n  \"gradient_clip_val\": None,\n  \"accumulate_grad_batches\": 1,\n  \"devices\": [0],\n}","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:25:04.527498Z","iopub.execute_input":"2024-04-23T04:25:04.527734Z","iopub.status.idle":"2024-04-23T04:25:04.576448Z","shell.execute_reply.started":"2024-04-23T04:25:04.527713Z","shell.execute_reply":"2024-04-23T04:25:04.575524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare Dataset","metadata":{}},{"cell_type":"code","source":"df = pl.read_parquet(\n    Path(data_dir, \"train.parquet\"),\n    columns=[\"molecule_smiles\", \"protein_name\", \"binds\"],\n    n_rows=N_ROWS,\n)\ntest_df = pl.read_parquet(\n    Path(data_dir, \"test.parquet\"),\n    columns=[\"molecule_smiles\"],\n    n_rows=10000 if DEBUG else None,\n)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:25:04.577621Z","iopub.execute_input":"2024-04-23T04:25:04.577955Z","iopub.status.idle":"2024-04-23T04:25:26.148615Z","shell.execute_reply.started":"2024-04-23T04:25:04.577920Z","shell.execute_reply":"2024-04-23T04:25:26.147080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfs = []\nfor i, protein_name in enumerate(PROTEIN_NAMES):\n    sub_df = df[i::3]\n    sub_df = sub_df.rename({\"binds\": protein_name})\n    if i == 0:\n        dfs.append(sub_df.drop([\"id\", \"protein_name\"]))\n    else:\n        dfs.append(sub_df[[protein_name]])\ndf = pl.concat(dfs, how=\"horizontal\")\ndf = df.sample(n=N_SAMPLES)\nprint(df.head())\nprint(df[PROTEIN_NAMES].sum())","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:25:26.150882Z","iopub.execute_input":"2024-04-23T04:25:26.151356Z","iopub.status.idle":"2024-04-23T04:25:30.695707Z","shell.execute_reply.started":"2024-04-23T04:25:26.151308Z","shell.execute_reply":"2024-04-23T04:25:30.694373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def normalize(x):\n    mol = Chem.MolFromSmiles(x)\n    smiles = Chem.MolToSmiles(mol, canonical=True, isomericSmiles=False)\n    return smiles\n\n\nif NORMALIZE:\n    df = df.with_columns(pl.col(\"molecule_smiles\").map_elements(normalize, return_dtype=pl.Utf8))\n    test_df = test_df.with_columns(pl.col(\"molecule_smiles\").map_elements(normalize, return_dtype=pl.Utf8))","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:25:45.397079Z","iopub.execute_input":"2024-04-23T04:25:45.397457Z","iopub.status.idle":"2024-04-23T04:25:56.554365Z","shell.execute_reply.started":"2024-04-23T04:25:45.397427Z","shell.execute_reply":"2024-04-23T04:25:56.553318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_idx, val_idx = train_test_split(np.arange(len(df)), test_size=0.2)\ntrain_df, val_df = df[train_idx], df[val_idx]\nlen(train_df), len(val_df)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:25:56.556369Z","iopub.execute_input":"2024-04-23T04:25:56.557089Z","iopub.status.idle":"2024-04-23T04:25:56.572576Z","shell.execute_reply.started":"2024-04-23T04:25:56.557053Z","shell.execute_reply":"2024-04-23T04:25:56.571714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Build Dataset","metadata":{}},{"cell_type":"code","source":"tokenizer = AutoTokenizer.from_pretrained(model_name)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:25:56.573763Z","iopub.execute_input":"2024-04-23T04:25:56.574117Z","iopub.status.idle":"2024-04-23T04:25:58.251181Z","shell.execute_reply.started":"2024-04-23T04:25:56.574084Z","shell.execute_reply":"2024-04-23T04:25:58.250422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def tokenize(batch, tokenizer):\n    output = tokenizer(batch[\"molecule_smiles\"], truncation=True)\n    return output\n\n\nclass LMDataset(Dataset):\n    def __init__(self, df, tokenizer, stage=\"train\"):\n        assert stage in [\"train\", \"val\", \"test\"]\n        self.tokenizer = tokenizer\n        self.stage = stage\n        df = (\n            datasets.Dataset\n            .from_pandas(df.to_pandas())\n            .map(tokenize, batched=True, fn_kwargs={\"tokenizer\": self.tokenizer})\n            .to_pandas()\n        )\n        self.df = pl.from_pandas(df)\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        data = self._generate_data(index)\n        data[\"label\"] = self._generate_label(index)\n        return data        \n\n    def _generate_data(self, index):\n        data = {\n            \"input_ids\": np.array(self.df[index, \"input_ids\"]),\n            \"attention_mask\": np.array(self.df[index, \"attention_mask\"]),\n        }\n        return data\n    \n    def _generate_label(self, index):\n        if self.stage == \"test\":\n            return np.array([0, 0, 0])\n        else:\n            return self.df[index, PROTEIN_NAMES].to_numpy()[0]\n\n\nLMDataset(train_df[:100], tokenizer)[0]","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:25:58.253277Z","iopub.execute_input":"2024-04-23T04:25:58.253574Z","iopub.status.idle":"2024-04-23T04:25:58.391622Z","shell.execute_reply.started":"2024-04-23T04:25:58.253549Z","shell.execute_reply":"2024-04-23T04:25:58.390766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LBDataModule(L.LightningDataModule):\n    def __init__(self, train_df, val_df, test_df, tokenizer):\n        super().__init__()\n        self.train_df = train_df\n        self.val_df = val_df\n        self.test_df = test_df\n        self.tokenizer = tokenizer\n\n    def _generate_dataset(self, stage):\n        if stage == \"train\":\n            df = self.train_df\n        elif stage == \"val\":\n            df = self.val_df\n        elif stage == \"test\":\n            df = self.test_df\n        else:\n            raise NotImplementedError\n        dataset = LMDataset(df, self.tokenizer, stage=stage)\n        return dataset\n\n    def _generate_dataloader(self, stage):\n        dataset = self._generate_dataset(stage)\n        if stage == \"train\":\n            shuffle=True\n            drop_last=True\n        else:\n            shuffle=False\n            drop_last=False\n        return DataLoader(\n            dataset,\n            batch_size=batch_size,\n            shuffle=shuffle,\n            drop_last=drop_last,\n            pin_memory=True,\n            collate_fn=DataCollatorWithPadding(self.tokenizer),\n        )\n\n    def train_dataloader(self):\n        return self._generate_dataloader(\"train\")\n\n    def val_dataloader(self):\n        return self._generate_dataloader(\"val\")\n\n    def test_dataloader(self):\n        return self._generate_dataloader(\"test\")\n    \n    \ndatamodule = LBDataModule(train_df, val_df, test_df, tokenizer)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:25:58.392778Z","iopub.execute_input":"2024-04-23T04:25:58.393295Z","iopub.status.idle":"2024-04-23T04:25:58.403132Z","shell.execute_reply.started":"2024-04-23T04:25:58.393267Z","shell.execute_reply":"2024-04-23T04:25:58.402272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Build Model","metadata":{}},{"cell_type":"code","source":"class LMModel(nn.Module):\n    def __init__(self, model_name):\n        super().__init__()\n        self.config = AutoConfig.from_pretrained(model_name, num_labels=3)\n        self.lm = AutoModel.from_pretrained(model_name, add_pooling_layer=False)\n        self.dropout = nn.Dropout(self.config.hidden_dropout_prob)\n        self.classifier = nn.Linear(self.config.hidden_size, self.config.num_labels)\n        self.loss_fn = nn.BCEWithLogitsLoss(reduction=\"mean\")\n\n    def forward(self, batch):\n        last_hidden_state = self.lm(\n            batch[\"input_ids\"],\n            attention_mask=batch[\"attention_mask\"],\n        ).last_hidden_state\n        logits = self.classifier(\n            self.dropout(last_hidden_state[:, 0])\n        )\n        return {\n            \"logits\": logits,\n        }\n\n    def calculate_loss(self, batch):\n        output = self.forward(batch)\n        loss = self.loss_fn(output[\"logits\"], batch[\"labels\"].float())\n        output[\"loss\"] = loss\n        return output\n\n    \nLMModel(model_name)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:25:58.404562Z","iopub.execute_input":"2024-04-23T04:25:58.405462Z","iopub.status.idle":"2024-04-23T04:26:00.282165Z","shell.execute_reply.started":"2024-04-23T04:25:58.405433Z","shell.execute_reply":"2024-04-23T04:26:00.281185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LBModelModule(L.LightningModule):\n    def __init__(self, model_name):\n        super().__init__()\n        self.model = LMModel(model_name)\n        self.map = AveragePrecision(task=\"binary\")\n\n    def forward(self, batch):\n        return self.model(batch)\n\n    def calculate_loss(self, batch, batch_idx):\n        return self.model.calculate_loss(batch)\n\n    def training_step(self, batch, batch_idx):\n        ret = self.calculate_loss(batch, batch_idx)\n        self.log(\"train_loss\", ret[\"loss\"], on_step=True, on_epoch=True, prog_bar=True, sync_dist=True)\n        return ret[\"loss\"]\n\n    def validation_step(self, batch, batch_idx):\n        ret = self.calculate_loss(batch, batch_idx)\n        self.log(\"val_loss\", ret[\"loss\"], on_step=False, on_epoch=True, prog_bar=True, sync_dist=True)\n        self.map.update(F.sigmoid(ret[\"logits\"]), batch[\"labels\"].long())\n\n    def on_validation_epoch_end(self):\n        val_map = self.map.compute()\n        self.log(\"val_map\", val_map, on_step=False, on_epoch=True, prog_bar=True, sync_dist=True)\n        self.map.reset()\n\n    def predict_step(self, batch, batch_idx, dataloader_idx=0):\n        logits = self.forward(batch)[\"logits\"]\n        probs = F.sigmoid(logits)\n        return probs\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=0.0001)\n        return {\n            \"optimizer\": optimizer,\n        }\n\n    \nmodelmodule = LBModelModule(model_name)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:26:00.283579Z","iopub.execute_input":"2024-04-23T04:26:00.283955Z","iopub.status.idle":"2024-04-23T04:26:00.503138Z","shell.execute_reply.started":"2024-04-23T04:26:00.283921Z","shell.execute_reply":"2024-04-23T04:26:00.502326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"checkpoint_callback = ModelCheckpoint(\n    filename=f\"model-{{val_map:.4f}}\",\n    save_weights_only=True,\n    monitor=\"val_map\",\n    mode=\"max\",\n    dirpath=\"/kaggle/working\",\n    save_top_k=1,\n    verbose=1,\n)\nearly_stop_callback = EarlyStopping(monitor=\"val_map\", mode=\"max\", patience=3)\nprogress_bar_callback = TQDMProgressBar(refresh_rate=1)\ncallbacks = [\n    checkpoint_callback,\n    early_stop_callback,\n    progress_bar_callback,\n]","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:26:00.504302Z","iopub.execute_input":"2024-04-23T04:26:00.504604Z","iopub.status.idle":"2024-04-23T04:26:00.519830Z","shell.execute_reply.started":"2024-04-23T04:26:00.504578Z","shell.execute_reply":"2024-04-23T04:26:00.518947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = L.Trainer(callbacks=callbacks, **trainer_params)\ntrainer.fit(modelmodule, datamodule)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:26:00.520963Z","iopub.execute_input":"2024-04-23T04:26:00.521240Z","iopub.status.idle":"2024-04-23T04:26:44.805119Z","shell.execute_reply.started":"2024-04-23T04:26:00.521217Z","shell.execute_reply":"2024-04-23T04:26:44.804095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"test_df = pl.read_parquet(\n    Path(data_dir, \"test.parquet\"),\n    columns=[\"molecule_smiles\"],\n    n_rows=10000 if DEBUG else None,\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:26:44.807748Z","iopub.execute_input":"2024-04-23T04:26:44.808060Z","iopub.status.idle":"2024-04-23T04:26:44.818284Z","shell.execute_reply.started":"2024-04-23T04:26:44.808034Z","shell.execute_reply":"2024-04-23T04:26:44.817359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"working_dir = Path(\"/kaggle/working\")\nmodel_paths = working_dir.glob(\"*.ckpt\")\ntest_dataloader = datamodule.test_dataloader()\nfor model_path in model_paths:\n    print(model_path)\n    modelmodule = LBModelModule.load_from_checkpoint(\n        checkpoint_path=model_path,\n        model_name=model_name,\n    )\n    predictions = trainer.predict(modelmodule, test_dataloader)\n    predictions = torch.cat(predictions).numpy()\n    pred_dfs = []\n    for i, protein_name in enumerate(PROTEIN_NAMES):\n        pred_dfs.append(\n            test_df.with_columns(\n                pl.lit(protein_name).alias(\"protein_name\"),\n                pl.lit(predictions[:, i]).alias(\"binds\"),\n            )\n        )\n    pred_df = pl.concat(pred_dfs)\n    submit_df = (\n        pl.read_parquet(Path(data_dir, \"test.parquet\"), columns=[\"id\", \"molecule_smiles\", \"protein_name\"])\n        .join(pred_df, on=[\"molecule_smiles\", \"protein_name\"], how=\"left\")\n        .select([\"id\", \"binds\"])\n        .sort(\"id\")\n    )\n    submit_df.write_csv(Path(working_dir, f\"submission_{model_path.stem}.csv\"))","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:26:44.819691Z","iopub.execute_input":"2024-04-23T04:26:44.820368Z","iopub.status.idle":"2024-04-23T04:26:50.404077Z","shell.execute_reply.started":"2024-04-23T04:26:44.820332Z","shell.execute_reply":"2024-04-23T04:26:50.403019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Ensemble","metadata":{}},{"cell_type":"code","source":"sub_files = list(working_dir.glob(\"submission_*.csv\"))\nsub_files","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:26:50.405360Z","iopub.execute_input":"2024-04-23T04:26:50.405668Z","iopub.status.idle":"2024-04-23T04:26:50.412173Z","shell.execute_reply.started":"2024-04-23T04:26:50.405642Z","shell.execute_reply":"2024-04-23T04:26:50.411351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_dfs = []\nfor sub_file in sub_files:\n    sub_dfs.append(pl.read_csv(sub_file))\nsubmit_df = (\n    pl.concat(sub_dfs)\n    .group_by(\"id\")\n    .agg(pl.col(\"binds\").mean())\n    .sort(\"id\")\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:26:50.413174Z","iopub.execute_input":"2024-04-23T04:26:50.413510Z","iopub.status.idle":"2024-04-23T04:26:50.729290Z","shell.execute_reply.started":"2024-04-23T04:26:50.413475Z","shell.execute_reply":"2024-04-23T04:26:50.728447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -fr *","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:26:58.793604Z","iopub.execute_input":"2024-04-23T04:26:58.793975Z","iopub.status.idle":"2024-04-23T04:27:00.381427Z","shell.execute_reply.started":"2024-04-23T04:26:58.793944Z","shell.execute_reply":"2024-04-23T04:27:00.380191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit_df.write_csv(Path(working_dir, \"submission.csv\"))","metadata":{"execution":{"iopub.status.busy":"2024-04-23T04:27:00.383862Z","iopub.execute_input":"2024-04-23T04:27:00.384685Z","iopub.status.idle":"2024-04-23T04:27:00.456370Z","shell.execute_reply.started":"2024-04-23T04:27:00.384642Z","shell.execute_reply":"2024-04-23T04:27:00.455474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Future Directions\n- Finding the optimal CV strategy\n- Increase data\n- Utilize buildingblock1 ~ buildingblock3\n- Large-scale models\n- Tune hyper-parameters\n- Ensemble\n- etc...","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}