{"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":67356,"databundleVersionId":8006601,"sourceType":"competition"},{"sourceId":170595844,"sourceType":"kernelVersion"},{"sourceId":181424750,"sourceType":"kernelVersion"},{"sourceId":181425964,"sourceType":"kernelVersion"}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Update of notebooks\n|version|update|cv|lb|\n|:--|:--|--:|--:|\n|v1|Molformer + LORA + Log sampling on 1.5M examples|-|0.334|\n|v3|v1 on 2M examples + canonical SMILES|-|0.310|\n|v4|v3 + save model|-|0.310|\n","metadata":{}},{"cell_type":"markdown","source":"# Install librairies","metadata":{}},{"cell_type":"code","source":"!pip install --quiet rdkit duckdb transformers torch datasets tqdm peft lightning","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:16:38.850367Z","iopub.execute_input":"2024-06-05T08:16:38.850721Z","iopub.status.idle":"2024-06-05T08:16:57.967923Z","shell.execute_reply.started":"2024-06-05T08:16:38.850688Z","shell.execute_reply":"2024-06-05T08:16:57.966819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define hyperparameters","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.insert(1, \"/kaggle/usr/lib\")\nfrom create_cv_indexes import get_cv_indexes\nfrom evaluate_utilities import calculate_ap, update_ap","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:16:57.969863Z","iopub.execute_input":"2024-06-05T08:16:57.970175Z","iopub.status.idle":"2024-06-05T08:17:17.545438Z","shell.execute_reply.started":"2024-06-05T08:16:57.970145Z","shell.execute_reply":"2024-06-05T08:17:17.544489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Number of training examples\nN_LIMIT = 1_800_000 # Maximum number of samples for training (will be N_LIMIT binding and N_LIMIT non binding)\nN_LIMIT = int(6*(N_LIMIT//6)) # Ensure that N_LIMIT can be divided by 6 to build dataset (3 proteins * 2 possible values for binding)\n\n# Tuning hyperparameters\nMAX_LENGTH = 380\nBATCH_SIZE = 164\n# BATCH_SIZE = 128\nLR = 1e-4\nNB_EPOCHS = 1\nGRADIENT_ACCUMULATION_STEPS = 3\n\n# LORA CONFIG\nLORA_CONFIG={\n    \"r\": 64, # Rank of the updated matrices\n    \"lora_alpha\": 128, # Scaling factor\n    \"lora_dropout\": 0.1,\n    \"bias\": \"none\",\n}\n# LORA_CONFIG = NONE # if one doesn't want to use the lora during the training and just train classification layers\n\nLOG_SAMPLING = True # If one wants to use log sampling during training\n\nPREDICT_PROTEINS = True # True if one wants to predict binding for all proteins at once, else will require protein_name\n\n# ECFPs\nINCLUDE_ECFPs = False # True if one wants to add ecfp representation \nBITS_ECFPs = 1024\nOUTPUTSHAPE_ECFP = 128\n\nIS_CANONICAL = True\nSAVING_PATH = \"bindingLM\"","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:17:17.546691Z","iopub.execute_input":"2024-06-05T08:17:17.547242Z","iopub.status.idle":"2024-06-05T08:17:17.554129Z","shell.execute_reply.started":"2024-06-05T08:17:17.547207Z","shell.execute_reply":"2024-06-05T08:17:17.553152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load data and librairies","metadata":{}},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nfrom typing import Literal, Optional\nfrom math import log\nimport warnings\nimport random\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\nimport gc\nimport os\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n\nimport torch\nfrom torch import nn\nfrom torch.nn import Linear, LeakyReLU, ReLU\nfrom torch.nn import functional as F\nfrom torch.utils.data import Dataset, DataLoader, Sampler\nfrom torch.optim import Adam\n\nfrom transformers import AutoModel, AutoTokenizer, AutoConfig, DataCollatorWithPadding, PreTrainedModel\nfrom datasets import Dataset, DatasetDict\n\nfrom tqdm import tqdm\n\nfrom transformers.utils.logging import disable_progress_bar\n\nfrom peft import LoraConfig, get_peft_model\n\nfrom rdkit import Chem\nfrom rdkit.Chem import AllChem\n\nimport lightning as L\n\ndisable_progress_bar()\nwarnings.filterwarnings(\"ignore\")\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-06-05T08:17:20.929928Z","iopub.execute_input":"2024-06-05T08:17:20.930775Z","iopub.status.idle":"2024-06-05T08:17:32.415052Z","shell.execute_reply.started":"2024-06-05T08:17:20.930742Z","shell.execute_reply":"2024-06-05T08:17:32.414247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import duckdb\ntrain_path = '/kaggle/input/belka-shrinking-the-dataset/train.parquet'\ntest_path = '/kaggle/input/leash-BELKA/test.parquet'\n\ncon = duckdb.connect()\n\n    \ndata = con.query(f\"\"\"SELECT *\n                        FROM parquet_scan('{train_path}')\n                        \"\"\").df()\n","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:17:32.416983Z","iopub.execute_input":"2024-06-05T08:17:32.417719Z","iopub.status.idle":"2024-06-05T08:18:09.401855Z","shell.execute_reply.started":"2024-06-05T08:17:32.417684Z","shell.execute_reply":"2024-06-05T08:18:09.401052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"proteins = [\"BRD4\",\"HSA\",\"sEH\"]\n\nmask = data[[f\"binds_{protein}\" for protein in proteins]].sum(axis=1) != 0\ndata[\"isbind\"] = mask\n\nn_total_samples = len(data)\nn_binds = data[\"isbind\"].sum()/n_total_samples\n\nbb_fq_no_binds = [data[~data[\"isbind\"]][f\"buildingblock{i+1}_smiles\"].value_counts(normalize=True).to_dict() for i in range(3)]\nbb_fq_binds = [data[data[\"isbind\"]][f\"buildingblock{i+1}_smiles\"].value_counts().to_dict() for i in range(3)]\nbinds_fq = [1-n_binds,n_binds]\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:18:09.402952Z","iopub.execute_input":"2024-06-05T08:18:09.403237Z","iopub.status.idle":"2024-06-05T08:18:51.512109Z","shell.execute_reply.started":"2024-06-05T08:18:09.403214Z","shell.execute_reply":"2024-06-05T08:18:51.511048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_index, index_eval_non_shared, index_random = get_cv_indexes(data)\n\ntrain_data = data.iloc[train_index].reset_index()\neval_random_data = data.iloc[index_random].reset_index()\neval_non_shared_data = data.iloc[index_eval_non_shared].reset_index()","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:18:51.514982Z","iopub.execute_input":"2024-06-05T08:18:51.515789Z","iopub.status.idle":"2024-06-05T08:19:52.856637Z","shell.execute_reply.started":"2024-06-05T08:18:51.515752Z","shell.execute_reply":"2024-06-05T08:19:52.855451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load and define model","metadata":{}},{"cell_type":"code","source":"model_name = \"ibm/MoLFormer-XL-both-10pct\"\n# model_name = \"DeepChem/ChemBERTa-10M-MTR\"\n# model_name = \"seyonec/ChemBERTa-zinc-base-v1\"\n\nmodel = AutoModel.from_pretrained(model_name, trust_remote_code=True)\ntokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True , max_length=MAX_LENGTH, padding=\"max_length\")","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:19:52.858042Z","iopub.execute_input":"2024-06-05T08:19:52.858354Z","iopub.status.idle":"2024-06-05T08:19:59.559083Z","shell.execute_reply.started":"2024-06-05T08:19:52.858327Z","shell.execute_reply":"2024-06-05T08:19:59.558340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BindingLM(PreTrainedModel):\n    def __init__(self, fm, n_proteins:int=3):\n        self.config = AutoConfig.from_pretrained(model_name,trust_remote_code=True, num_labels=n_proteins)\n        \n        super().__init__(self.config)\n        self.DEVICE = (\n        torch.device(\"cuda\")\n        if torch.cuda.is_available()\n        else (\n            torch.device(\"mps\")\n            if torch.backends.mps.is_available()\n            else torch.device(\"cpu\")\n        )\n    )\n        self.n_proteins = n_proteins\n        \n        self.fm = fm\n        self.loss_fn = nn.BCELoss(reduction=\"mean\")\n        \n        self.dropout = nn.Dropout(0.1)\n        \n    \n        n = self.config.hidden_size \n        \n        self.c1 = nn.Linear(n, n//2)\n        \n        if INCLUDE_ECFPs:\n            self.c2 = nn.Linear(OUTPUTSHAPE_ECFP+n//2, n//4)\n            self.c3 = nn.Linear(n//4, n//8)\n        else:\n            self.c2 = nn.Linear(n//2, n//4)\n            self.c3 = nn.Linear(n//4, n//8)\n            \n        self.clf = nn.Linear(n//8, self.config.num_labels)\n        \n        for param in self.fm.parameters():\n            param.requires_grad_(False) # freeze the model - train adapters later\n        \n        if INCLUDE_ECFPs:\n            self.conv1 = nn.Conv1d(in_channels=1, out_channels=32, kernel_size=5, stride=2)\n            self.conv2 = nn.Conv1d(in_channels=32, out_channels=64, kernel_size=5, stride=2)\n            self.pool = nn.MaxPool1d(2)\n            self.ecfps_lin_shape = 64 * 63\n            self.fc1 = nn.Linear(self.ecfps_lin_shape, 128)\n            self.fc2 = nn.Linear(128, OUTPUTSHAPE_ECFP)\n        \n        self.activation = nn.ReLU()\n        \n        ## Define LORA config\n        if LORA_CONFIG is not None:\n            # cf. to choose target modules for LORA ft: https://huggingface.co/docs/peft/developer_guides/custom_models\n            \n            target_modules = []\n            for layer_name, layer_type in [(n, type(m)) for n, m in self.fm.named_modules()]:\n                if layer_type in [torch.nn.modules.linear.Linear,torch.nn.modules.sparse.Embedding]:\n                    target_modules.append(layer_name)\n\n            lora_config = LoraConfig(\n            **LORA_CONFIG,\n            target_modules=target_modules,\n        )\n        \n            self.fm = get_peft_model(self.fm, lora_config)\n        self.to(self.DEVICE)\n\n    \n    def forward(self, batch):\n        input_ids = batch[\"input_ids\"].to(self.DEVICE)\n        attention_mask = batch[\"attention_mask\"].to(self.DEVICE)\n        \n        if INCLUDE_ECFPs:\n            ecfps = batch[\"ecfps\"].to(self.DEVICE)\n            ecfps = ecfps.unsqueeze(1)\n            ecfps = self.conv1(ecfps)\n            ecfps = self.pool(ecfps)\n            ecfps = self.conv2(ecfps)\n            ecfps = self.pool(ecfps)\n            ecfps = ecfps.view(-1, self.ecfps_lin_shape)\n            ecfps = self.fc1(ecfps)\n            ecfps = self.fc2(ecfps)\n        \n        x = self.fm(input_ids, attention_mask).last_hidden_state[:, 0]\n\n        x = self.c1(x)\n        x = self.dropout(x)\n        x = self.activation(x)\n        \n        if INCLUDE_ECFPs:\n            x = torch.cat([x, ecfps], dim=1)\n\n        x = self.c2(x)\n        x = self.dropout(x)\n        x = self.activation(x)\n        \n        x = self.c3(x)\n        x = self.dropout(x)\n        x = self.activation(x)\n        \n        x = self.clf(x)\n        x = nn.Sigmoid()(x)\n        return x\n    \nbindingLm = BindingLM(model)\ndevice = bindingLm.DEVICE","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:19:59.560506Z","iopub.execute_input":"2024-06-05T08:19:59.560804Z","iopub.status.idle":"2024-06-05T08:20:00.671228Z","shell.execute_reply.started":"2024-06-05T08:19:59.560776Z","shell.execute_reply":"2024-06-05T08:20:00.670141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define dataset","metadata":{}},{"cell_type":"code","source":"# Generate ECFPs\n# Thanks to https://www.kaggle.com/code/andrewdblevins/leash-tutorial-ecfps-and-random-forest\ndef generate_ecfp(molecule, radius: int = 2, bits: int =BITS_ECFPs):\n    molecule = Chem.MolFromSmiles(molecule)\n    if molecule is None:\n        return None\n    return list(AllChem.GetMorganFingerprintAsBitVect(molecule, radius, nBits=bits))\n\n\n\nclass BindingDataset(Dataset):\n    def __init__(self, dataframe: pd.DataFrame, test: bool = False):\n        self.dataframe = dataframe\n        self.test = test\n        self.bind_cols = [f\"binds_{prot}\" for prot in proteins]\n        self.cols = [\"molecule_smiles\"]\n        if not self.test:\n            self.cols = self.cols + self.bind_cols\n        \n    def __getitem__(self, idx):\n        rows = self.dataframe.iloc[idx]\n        \n        if INCLUDE_ECFPs:\n            self.cols.append(\"ecfps\")\n            rows.loc[:,\"ecfps\"] = rows[\"molecule_smiles\"].apply(generate_ecfp)\n            \n        rows = rows[self.cols]\n        if IS_CANONICAL:\n            rows.loc[:,\"molecule_smiles\"] = rows.loc[:,\"molecule_smiles\"].apply(lambda x: Chem.MolToSmiles(Chem.MolFromSmiles(x), canonical=True) )\n        rows.loc[:,\"molecule_smiles\"] = rows.loc[:,\"molecule_smiles\"].apply(lambda x: tokenizer(x))\n        rows.loc[:, \"input_ids\"] = rows.loc[:,\"molecule_smiles\"].apply(lambda x: x[\"input_ids\"])\n        rows.loc[:, \"attention_mask\"] = rows.loc[:,\"molecule_smiles\"].apply(lambda x: x[\"attention_mask\"])\n        rows = rows.drop(columns=[\"molecule_smiles\"])\n        if not self.test:\n            rows[\"binds\"] = rows.apply(lambda x: [x[col] for col in self.bind_cols], axis=1)\n            rows = rows.drop(columns=self.bind_cols)\n            \n        for column in rows.columns:\n            if column == \"input_ids\":\n                rows[column] = rows[column].apply(lambda x: torch.IntTensor(np.array(x)))\n            else:\n                rows[column] = rows[column].apply(lambda x: torch.FloatTensor(np.array(x).astype(np.float32)))\n            \n        rows = rows.reset_index()\n        try:\n            rows.drop(columns=[\"id\"])\n        except:\n            pass\n        return rows\n\n    def __len__(self):\n        return len(self.dataframe)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:21:08.543380Z","iopub.execute_input":"2024-06-05T08:21:08.544080Z","iopub.status.idle":"2024-06-05T08:21:08.559478Z","shell.execute_reply.started":"2024-06-05T08:21:08.544043Z","shell.execute_reply":"2024-06-05T08:21:08.558473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define data sampler","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nfrom concurrent.futures import ThreadPoolExecutor\n\nclass LogSampler(Sampler):\n    def __init__(\n        self,\n        dataset: BindingDataset,\n        bb_fq_no_binds: Optional[list[dict]]=None,\n        bb_fq_binds: Optional[list[dict]]=None,\n        binds_fq: Optional[list[float]]=None,\n        batch_size: int = 512,\n    ) -> None:\n        self.dataset = dataset\n        self.batch_size = batch_size\n        \n        self.bb_fq_no_binds = bb_fq_no_binds\n        self.bb_fq_binds = bb_fq_binds\n        self.binds_fq = binds_fq\n\n        n_total_samples = len(self.dataset.dataframe)\n\n        # Precompute frequencies\n        for ind, (elem1, elem2) in enumerate(zip(self.bb_fq_no_binds, self.bb_fq_binds)):\n            for ind2, elem in enumerate([elem1, elem2]):\n                N = n_total_samples * self.binds_fq[ind2]\n\n                fqs = {nam: np.log(1 + N * fq) for nam, fq in elem.items()}\n                sum_fqs = sum(fqs.values())\n                fqs = {nam: fq / sum_fqs for nam, fq in fqs.items()}\n\n                if ind2 == 0:\n                    self.bb_fq_no_binds[ind] = fqs\n                else:\n                    self.bb_fq_binds[ind] = fqs\n\n        # Precompute frequencies for binding\n        self.binds_fq = [np.log(1 + len(data) * fq) for fq in self.binds_fq]\n        self.binds_fq = np.array(self.binds_fq) / np.sum(self.binds_fq)\n        \n        self.nb_batches = N_LIMIT // self.batch_size\n        \n        self.on_epoch_end()\n\n    def __iter__(self):\n        \n        while self.batches_on <= self.nb_batches:\n            # For the log sampling, to keep it efficient, every 300 steps we sample 5% of total dataset to work on\n            if self.batches_on % 300 == 0 or self.batches_on == 0:\n                sample_batches = self.dataset.dataframe.sample(int(0.05 * len(self.dataset.dataframe)))\n            self.batches_on += 1\n            \n            mask_bbi = np.random.choice([1, 2, 3], size=self.batch_size)\n            nbs_bbis = np.array([np.sum(mask_bbi == i) for i in range(1, 4)])\n            list_samples = []\n            \n            binds_mask = np.random.choice([0, 1], p=self.binds_fq, size=self.batch_size)\n            nb_binds = np.array([np.sum(binds_mask == i) for i in range(2)])\n\n            # Define a function for sampling data_bind\n            def sample_data_bind(bind, bbi, frequencies, nb_bind):\n                data_bind = sample_batches[sample_batches[\"isbind\"] == bind]\n                k = int((nb_bind * nbs_bbis[bbi]) / (self.batch_size * 3))\n                return data_bind.sample(n=k, weights=data_bind[f\"buildingblock{bbi+1}_smiles\"].map(frequencies))\n            \n            # Use ThreadPoolExecutor for parallel sampling\n            with ThreadPoolExecutor() as executor:\n                samples_per_bind = list(executor.map(sample_data_bind, [0, 1], range(3), self.bb_fq_no_binds, nb_binds))\n\n            samples = pd.concat(samples_per_bind, axis=0)\n            \n            if len(samples) < self.batch_size:\n                add_samples = sample_batches.sample(self.batch_size - len(samples))\n                samples = pd.concat([samples, add_samples], axis=0)\n\n            samples = samples.sample(frac=1)\n            yield samples.index\n\n    def on_epoch_end(self) -> None:\n        self.batches_on = 0\n        self.indices = self.dataset.dataframe.index\n\n    def __len__(self) -> int:\n        return int(np.ceil(len(self.dataset.dataframe) / self.batch_size))\n","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:21:08.879353Z","iopub.execute_input":"2024-06-05T08:21:08.880000Z","iopub.status.idle":"2024-06-05T08:21:08.899393Z","shell.execute_reply.started":"2024-06-05T08:21:08.879970Z","shell.execute_reply":"2024-06-05T08:21:08.898451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define data loader","metadata":{}},{"cell_type":"code","source":"def prepare_loader(dataset: pd.DataFrame, shuffle: bool = False, test: bool=False, log_sampling: bool= False):\n    dataset = BindingDataset(dataset, test)\n    if log_sampling:\n        sampler = LogSampler(dataset, bb_fq_no_binds, bb_fq_binds, binds_fq, batch_size=BATCH_SIZE)\n        # Not enough RAM to make \n        # dataloader = DataLoader(dataset, collate_fn=DataCollatorWithPadding(tokenizer), batch_sampler=sampler, num_workers=2)\n        dataloader = DataLoader(dataset, collate_fn=DataCollatorWithPadding(tokenizer), batch_sampler=sampler)\n    else:\n        sampler = None\n        batch_size = BATCH_SIZE\n        dataloader = DataLoader(dataset, collate_fn=DataCollatorWithPadding(tokenizer), batch_size=BATCH_SIZE)\n        # dataloader = DataLoader(dataset, collate_fn=DataCollatorWithPadding(tokenizer), batch_size=BATCH_SIZE, num_workers=2)\n    return dataloader\n\ntrain_dataloader = prepare_loader(train_data, True, log_sampling=LOG_SAMPLING)\neval_random_dataloader = prepare_loader(eval_random_data)\neval_non_shared_dataloader = prepare_loader(eval_non_shared_data)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:21:09.359061Z","iopub.execute_input":"2024-06-05T08:21:09.359956Z","iopub.status.idle":"2024-06-05T08:21:09.375628Z","shell.execute_reply.started":"2024-06-05T08:21:09.359919Z","shell.execute_reply":"2024-06-05T08:21:09.374772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training loop","metadata":{}},{"cell_type":"code","source":"gc.collect()\ntorch.cuda.empty_cache()\n\noptim = Adam(\n            params=bindingLm.parameters(),\n            lr=LR,\n            betas=(0.5, 0.9),\n        )\nn_batchs = N_LIMIT // BATCH_SIZE\nsteps_eval = 1000\n\ntrain_scores = []\nrandom_split_scores = []\nnon_shared_split_scores = []\n\nfor epoch in range(NB_EPOCHS):\n    train_preds = []\n    train_targets = []\n    \n    bindingLm.train()\n    with tqdm(train_dataloader, total=n_batchs) as tepoch:\n        tepoch.set_description(f\"Epoch {epoch+1}/{NB_EPOCHS}: Train set\")\n        for ind,batch in enumerate(tepoch):\n            targets = batch[\"binds\"]\n            targets = targets.type(torch.FloatTensor)\n            targets = targets.to(device)\n            \n            predictions = bindingLm.forward(batch)\n            loss = bindingLm.loss_fn(predictions, targets)\n            loss.backward()\n            \n            if (ind%steps_eval ==0) and (ind//steps_eval>=1) :\n                score = calculate_ap(torch.cat(train_preds, dim=0), torch.cat(train_targets, dim=0), proteins)\n                train_scores.append(score)\n                train_preds = []\n                train_targets = []\n            else:\n                train_preds.append(predictions.cpu().detach())\n                train_targets.append(targets.cpu().detach())\n            \n            \n            if ind%GRADIENT_ACCUMULATION_STEPS==0:\n                optim.step()\n                optim.zero_grad()\n            \n            tepoch.set_postfix(batch=f\"{ind+1}/{n_batchs}\",loss = \"{0:.3f}\".format(loss.cpu().detach().item()))\n            \n            # Release some GPU memory\n            gc.collect()\n            torch.cuda.empty_cache()\n    del train_preds\n    del train_targets\n           \n    bindingLm.eval()\n    # Evaluate random_split_score\n    with tqdm(eval_random_dataloader) as randomeval:\n        random_split_preds = []\n        random_split_targets = []\n        for ind,batch in enumerate(randomeval):\n            predictions = bindingLm.forward(batch)\n            targets = batch[\"binds\"]\n            targets = targets.type(torch.FloatTensor)\n            \n            random_split_preds.append(predictions.cpu().detach())\n            random_split_targets.append(targets)\n            \n            gc.collect()\n            torch.cuda.empty_cache()\n        score = calculate_ap(torch.cat(random_split_preds, dim=0), torch.cat(random_split_targets, dim=0), proteins)\n        random_split_scores.append(score)\n        \n    del random_split_preds\n    del random_split_targets\n    \n    \n    # Evaluate non shared bb split\n    with tqdm(eval_non_shared_dataloader) as nnsharedeval:\n        non_shared_bb_preds = []\n        non_shared_bb_targets = []\n        \n        for ind,batch in enumerate(nnsharedeval):\n            predictions = bindingLm.forward(batch)\n            targets = batch[\"binds\"]\n            targets = targets.type(torch.FloatTensor)\n            \n            non_shared_bb_preds.append(predictions.cpu().detach())\n            non_shared_bb_targets.append(targets)\n            \n            \n            gc.collect()\n            torch.cuda.empty_cache()\n            \n        score = calculate_ap(torch.cat(non_shared_bb_preds, dim=0), torch.cat(non_shared_bb_targets, dim=0), proteins)\n        non_shared_split_scores.append(score)\n    \n    del non_shared_bb_preds\n    del non_shared_bb_targets","metadata":{"execution":{"iopub.status.busy":"2024-06-05T08:21:10.024488Z","iopub.execute_input":"2024-06-05T08:21:10.024858Z","iopub.status.idle":"2024-06-05T08:54:22.975229Z","shell.execute_reply.started":"2024-06-05T08:21:10.024826Z","shell.execute_reply":"2024-06-05T08:54:22.974164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Display CV results","metadata":{}},{"cell_type":"code","source":"print(\"----------------------------------------\")\nprint(\"Train scores: \\n\")\nfor key, item in train_scores[-1].items():\n    scr = '{0:.4f}'.format(item)\n    print(f\"{key}: {scr}\")\n    \n    \nprint(\"\\n\\n----------------------------------------\")\nprint(\"Random split scores: \\n\")\nfor key, item in random_split_scores[-1].items():\n    scr = '{0:.4f}'.format(item)\n    print(f\"{key}: {scr}\")\n    \n    \nprint(\"\\n\\n----------------------------------------\")\nprint(\"Non shared bbs scores: \\n\")\nfor key, item in non_shared_split_scores[-1].items():\n    scr = '{0:.4f}'.format(item)\n    print(f\"{key}: {scr}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-06-05T09:00:21.052143Z","iopub.execute_input":"2024-06-05T09:00:21.052963Z","iopub.status.idle":"2024-06-05T09:00:21.059394Z","shell.execute_reply.started":"2024-06-05T09:00:21.052933Z","shell.execute_reply":"2024-06-05T09:00:21.058417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define lightning module","metadata":{}},{"cell_type":"code","source":"# define the LightningModule\nclass LitBindingLM(L.LightningModule):\n    def __init__(self, bindingLm: BindingLM):\n        super().__init__()\n        self.bindingLm = bindingLm\n\n    def training_step(self, batch, batch_idx):\n        targets = batch[\"binds\"]\n        targets = targets.type(torch.FloatTensor)\n        predictions = self.bindingLm.forward(batch)\n        loss = self.bindingLm.loss_fn(predictions, targets)\n        return loss\n\n    def configure_optimizers(self):\n        optimizer = Adam(self.bindingLm.parameters(), lr=LR)\n        return optimizer\n    \nlitbindinglm = LitBindingLM(bindingLm)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\ntrainer = L.Trainer(accelerator=\"gpu\", devices=1, limit_train_batches=100, max_epochs=NB_EPOCHS, default_root_dir=\"/kaggle/working/models\")\ntrainer.fit(model=litbindinglm, train_dataloaders=dataloader)\n'''\n# Raises an error for now (that is weird consired that normal training loop doesn't)\n# RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cuda:0 and cpu!\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Making some free space\n# del dataloader\nsubmission = pd.read_parquet(test_path)\n\nprediction_submission = submission[['buildingblock1_smiles', 'buildingblock2_smiles', 'buildingblock3_smiles', 'molecule_smiles']].drop_duplicates()\ntest_loader = prepare_loader(prediction_submission, test=True)\n    \nn_batchs = len(submission) // BATCH_SIZE\n\nbinds_predictions = []\n\n\nwith tqdm(test_loader, total=n_batchs) as ttest:\n    bindingLm.eval()\n    for ind,batch in enumerate(test_loader):\n        ttest.set_postfix(batch=f\"{ind+1}/{n_batchs}\")\n        with torch.no_grad():\n            predictions = bindingLm.forward(batch)\n            binds_predictions.append(predictions.cpu().detach())\n        torch.cuda.empty_cache()\n        \npredictions = torch.cat(binds_predictions, dim=0)\n\n    \n\npredictions = list(predictions.numpy())\npredictions = [predictions[i] if i<len(predictions) else np.array([0.0,0.0,0.0]) for i in range(len(prediction_submission))]\nprediction_submission[\"binds\"] = predictions\nprediction_submission = prediction_submission.explode('binds')\nprediction_submission['protein_name'] = [proteins[i % len(proteins)] for i in range(len(prediction_submission))]\nsubmission = submission.merge(prediction_submission, how=\"left\")\n    \n        \nsubmission[['id', 'binds']].to_csv('submission.csv', index = False)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Save model","metadata":{}},{"cell_type":"code","source":"bindingLm.save_pretrained(SAVING_PATH,safe_serialization=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}