{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":67356,"databundleVersionId":8006601,"sourceType":"competition"},{"sourceId":7856946,"sourceType":"datasetVersion","datasetId":4608382}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<h1>Putting on my chemist robe and glasses :)</h1><center><img src=\"https://static.displate.com/280x392/displate/2022-12-28/494bd2c86cb966c2129375457a82982f_80fd640e09154845b9ca853fda67a96d.jpg\" ></center>Please, check this recent paper, <a href=\"https://arxiv.org/abs/2402.09391\">LlaSMol: Advancing Large Language Models for Chemistry with a Large-Scale, Comprehensive, High-Quality Instruction Tuning Dataset</a> to understand which model I am using.  \n<br><br>So, let us first figure out what they want from us and what we want from our code. <br><br>\nThe dataset comprises binary classification data representing whether small molecules bind to three different protein targets. Each entry includes SMILES representations of molecule structures and binary labels for binding to each protein target. The competition data, provided by Leash Biosciences, consists of approximately **98M** training examples per protein, **200K** validation examples per protein, and **360K** test molecules per protein. *Keep in mind that dataset is very imbalanced:  roughly 0.5% of examples are classified as binders*<br><br>\n\n<center><h2>So, what is SMILES?</h2><img src=\"https://static.wixstatic.com/media/7cced3_08c22a1ad56a4d2b81c46c8d71ebc34b~mv2.gif\" ></center><br>\n\nSMILES (Simplified Molecular Input Line Entry System) is a notation system for representing chemical structures in a computer-readable format. Developed with funding from the U.S. Environmental Protection Agency, it offers a flexible and easily learned approach to representing molecules. SMILES notation follows five basic syntax rules:\n\n* Atoms and Bonds: Atoms are represented by their atomic symbols, with lowercase letters indicating aromatic atoms. Bonds are denoted by symbols (- for single, = for double, # for triple, * for aromatic, and . for disconnected structures).\n* Simple Chains: Chains of atoms are represented by combining atomic symbols and bond symbols. Hydrogen atoms are suppressed unless explicitly stated.\n* Branches: Branches from chains are enclosed in parentheses and placed directly after the atom to which they are connected.\n* Rings: Ring structures are identified by using numbers to denote the opening and closing ring atoms. Different numbers are used for each ring, and bond symbols may precede the ring closure number.\n* Charged Atoms: Charges on atoms are indicated by placing the atom symbol within brackets, enclosing the charge.<br>\n\n\n<center><h2>What are the targets?</h2></center><br>\n\nThe dataset includes three protein targets:\n<h3>EPHX2 (sEH):</h3><center><img src=\"https://www.researchgate.net/publication/26836510/figure/fig1/AS:394310370512898@1471022326590/Pathways-of-EETs-synthesis-metabolism-and-action.png\" ></center><br>\n* This target refers to epoxide hydrolase 2, encoded by the EPHX2 genetic locus. Its protein product, commonly known as soluble epoxide hydrolase (sEH), is an enzyme that catalyzes certain chemical reactions and hydrolyzes phosphate groups. It is a potential drug target for conditions like high blood pressure and diabetes. The dataset includes screening data obtained from Leash Biosciences, along with structural information for model evaluation.\n<h3>BRD4:</h3><center><img src=\"https://ars.els-cdn.com/content/image/1-s2.0-S1043661823001238-ga1.jpg\" ></center><br>\n* Bromodomain 4 is encoded by the BRD4 locus. Its protein product, also named BRD4, plays a role in gene transcription regulation by binding to histones in the nucleus. This protein is implicated in cancer progression, and inhibiting its activity has been explored as a therapeutic strategy. The dataset includes screening data from Leash Biosciences and structural information for model evaluation.\n<h3>ALB (HSA):</h3><center><img src=\"https://ars.els-cdn.com/content/image/1-s2.0-S0162013419301734-gr4.jpg\" ></center><br>\n* Serum albumin, encoded by the ALB locus, is the most abundant protein in blood. It regulates osmotic pressure and transports various molecules, including drugs, hormones, and fatty acids. Predicting the binding of small molecules to albumin is crucial for drug development, as it impacts drug distribution and effectiveness. The dataset includes screening data from Leash Biosciences and structural information for model evaluation.\n","metadata":{}},{"cell_type":"markdown","source":"<h3>Now let's get the things done! :D</h3> <br>\n\nWe shall start with the installation of necessary tools. <br><br>\n**PyTorch Lightning is a Python library that simplifies training deep learning models by providing a standardized interface and automating common tasks like logging and reproducibility. It enables faster model development and training by abstracting away boilerplate code and supporting distributed training across multiple GPUs and machines..**","metadata":{}},{"cell_type":"code","source":"!pip install -U /kaggle/input/lightning-2-2-1/lightning-2.2.1-py3-none-any.whl","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-04-23T19:39:59.988150Z","iopub.execute_input":"2024-04-23T19:39:59.988543Z","iopub.status.idle":"2024-04-23T19:40:15.834170Z","shell.execute_reply.started":"2024-04-23T19:39:59.988513Z","shell.execute_reply":"2024-04-23T19:40:15.833082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**RDKit is a widely-used open-source toolkit for cheminformatics. It offers various functions for working with chemical structures and data, including molecular representation, substructure searching, and compound library handling. RDKit is popular in pharmaceutical research and drug discovery due to its comprehensive features and easy integration with Python.**","metadata":{}},{"cell_type":"code","source":"!pip install rdkit\n!pip install peft","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-04-23T19:40:24.182858Z","iopub.execute_input":"2024-04-23T19:40:24.183274Z","iopub.status.idle":"2024-04-23T19:40:51.654314Z","shell.execute_reply.started":"2024-04-23T19:40:24.183240Z","shell.execute_reply":"2024-04-23T19:40:51.653056Z"},"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, AutoModelForCausalLM\nimport datasets\nfrom rdkit import Chem\nfrom peft import PeftModelForCausalLM, PeftModel","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-23T19:43:15.605870Z","iopub.execute_input":"2024-04-23T19:43:15.609846Z","iopub.status.idle":"2024-04-23T19:43:15.617766Z","shell.execute_reply.started":"2024-04-23T19:43:15.609777Z","shell.execute_reply":"2024-04-23T19:43:15.616750Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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 = \"osunlp/LlaSMol-Mistral-7B\"\nbase_model = \"mistralai/Mistral-7B-v0.1\"\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-23T19:41:18.195223Z","iopub.execute_input":"2024-04-23T19:41:18.195935Z","iopub.status.idle":"2024-04-23T19:41:18.202921Z","shell.execute_reply.started":"2024-04-23T19:41:18.195897Z","shell.execute_reply":"2024-04-23T19:41:18.201827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Further, we convert SMILES representations of molecules into RDKit molecule objects and then generates ECFP fingerprints for each molecule, storing the results in the DataFrame. \n\n**ECFP, or Extended Connectivity Fingerprints**, are molecular fingerprints used in cheminformatics to encode structural features of molecules. They represent connectivity patterns of atoms by hashing local environments within a specified radius. ECFP fingerprints are widely used for similarity searching and activity prediction in drug discovery and cheminformatics due to their efficiency and robustness.\n\n<center><img src=\"https://miro.medium.com/max/652/1*YjmaYuWldV2yFH1TvPvNNw.jpeg\" ></center><br>","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-23T19:41:24.395027Z","iopub.execute_input":"2024-04-23T19:41:24.395953Z","iopub.status.idle":"2024-04-23T19:41:44.422549Z","shell.execute_reply.started":"2024-04-23T19:41:24.395917Z","shell.execute_reply":"2024-04-23T19:41:44.421334Z"},"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-23T19:41:47.765504Z","iopub.execute_input":"2024-04-23T19:41:47.766068Z","iopub.status.idle":"2024-04-23T19:41:52.812631Z","shell.execute_reply.started":"2024-04-23T19:41:47.766026Z","shell.execute_reply":"2024-04-23T19:41:52.811633Z"},"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-23T18:16:47.386977Z","iopub.execute_input":"2024-04-23T18:16:47.387339Z","iopub.status.idle":"2024-04-23T18:54:13.417345Z","shell.execute_reply.started":"2024-04-23T18:16:47.387311Z","shell.execute_reply":"2024-04-23T18:54:13.416452Z"},"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-23T19:42:03.098753Z","iopub.execute_input":"2024-04-23T19:42:03.099497Z","iopub.status.idle":"2024-04-23T19:42:03.333730Z","shell.execute_reply.started":"2024-04-23T19:42:03.099462Z","shell.execute_reply":"2024-04-23T19:42:03.332801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\naccess_token = user_secrets.get_secret(\"Kaggle\")","metadata":{"execution":{"iopub.status.busy":"2024-04-23T19:42:05.462382Z","iopub.execute_input":"2024-04-23T19:42:05.462761Z","iopub.status.idle":"2024-04-23T19:42:05.624529Z","shell.execute_reply.started":"2024-04-23T19:42:05.462735Z","shell.execute_reply":"2024-04-23T19:42:05.623688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenizer = AutoTokenizer.from_pretrained(base_model, token=access_token)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T19:42:06.913906Z","iopub.execute_input":"2024-04-23T19:42:06.914278Z","iopub.status.idle":"2024-04-23T19:42:07.532704Z","shell.execute_reply.started":"2024-04-23T19:42:06.914249Z","shell.execute_reply":"2024-04-23T19:42:07.531514Z"},"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-23T19:42:18.617503Z","iopub.execute_input":"2024-04-23T19:42:18.617942Z","iopub.status.idle":"2024-04-23T19:42:18.799115Z","shell.execute_reply.started":"2024-04-23T19:42:18.617909Z","shell.execute_reply":"2024-04-23T19:42:18.798000Z"},"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":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-04-23T19:42:24.808695Z","iopub.execute_input":"2024-04-23T19:42:24.809580Z","iopub.status.idle":"2024-04-23T19:42:24.820821Z","shell.execute_reply.started":"2024-04-23T19:42:24.809534Z","shell.execute_reply":"2024-04-23T19:42:24.819843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2024-04-23T19:42:28.178769Z","iopub.execute_input":"2024-04-23T19:42:28.179710Z","iopub.status.idle":"2024-04-23T19:42:28.219272Z","shell.execute_reply.started":"2024-04-23T19:42:28.179672Z","shell.execute_reply":"2024-04-23T19:42:28.218447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LMModel(nn.Module):\n    def __init__(self, model_name):\n        super().__init__()\n        self.config = AutoConfig.from_pretrained(base_model, num_labels=3, token=access_token)\n        self.lm = AutoModel.from_pretrained(\n            base_model,\n            #add_pooling_layer=False,\n            #torch_dtype=torch.bfloat16,\n            #device_map=\"auto\",\n            token=access_token\n        )   \n        self.lm = PeftModel.from_pretrained(\n            self.lm,\n            model_name,\n            torch_dtype=torch.bfloat16,\n            offload_dir = \"/kaggle/working/\"\n        )\n        self.lm = self.lm.merge_and_unload()\n        self.dropout = nn.Dropout(self.config.hidden_dropout_prob)\n        self.classifier = nn.Linear(self.config.hidden_size, 3)\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","metadata":{"execution":{"iopub.status.busy":"2024-04-23T19:45:18.148078Z","iopub.execute_input":"2024-04-23T19:45:18.148483Z","iopub.status.idle":"2024-04-23T19:45:18.158784Z","shell.execute_reply.started":"2024-04-23T19:45:18.148451Z","shell.execute_reply":"2024-04-23T19:45:18.157720Z"},"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-23T19:45:19.452245Z","iopub.execute_input":"2024-04-23T19:45:19.452649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-23T19:15:07.956394Z","iopub.execute_input":"2024-04-23T19:15:07.956738Z","iopub.status.idle":"2024-04-23T19:15:07.965080Z","shell.execute_reply.started":"2024-04-23T19:15:07.956712Z","shell.execute_reply":"2024-04-23T19:15:07.964080Z"},"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-23T19:15:10.497354Z","iopub.execute_input":"2024-04-23T19:15:10.497762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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_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_count":null,"outputs":[]},{"cell_type":"code","source":"sub_files = list(working_dir.glob(\"submission_*.csv\"))\nsub_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_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -fr *","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit_df.write_csv(Path(working_dir, \"submission.csv\"))","metadata":{},"execution_count":null,"outputs":[]}]}