{"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"}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\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\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\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-04-28T16:03:25.147116Z","iopub.execute_input":"2024-04-28T16:03:25.147512Z","iopub.status.idle":"2024-04-28T16:03:26.083454Z","shell.execute_reply.started":"2024-04-28T16:03:25.147459Z","shell.execute_reply":"2024-04-28T16:03:26.082148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PROTEIN_NAME = 'BRD4'\nBASE_MODEL = 'DeepChem/ChemBERTa-10M-MTR'\nTRAIN_RAW_DIR = \"/kaggle/input/leash-BELKA/train.parquet\"\nTEST_RAW_DIR = \"/kaggle/input/leash-BELKA/test.parquet\"","metadata":{"execution":{"iopub.status.busy":"2024-04-28T16:03:26.085768Z","iopub.execute_input":"2024-04-28T16:03:26.086187Z","iopub.status.idle":"2024-04-28T16:03:26.091234Z","shell.execute_reply.started":"2024-04-28T16:03:26.08616Z","shell.execute_reply":"2024-04-28T16:03:26.090072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install torch transformers evaluate","metadata":{"execution":{"iopub.status.busy":"2024-04-28T16:03:26.092623Z","iopub.execute_input":"2024-04-28T16:03:26.092976Z","iopub.status.idle":"2024-04-28T16:03:41.891649Z","shell.execute_reply.started":"2024-04-28T16:03:26.09295Z","shell.execute_reply":"2024-04-28T16:03:41.890306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_data(raw_data_dir, write_data_dir, protein_name=PROTEIN_NAME):\n    # filters for only the relevant protein\n    # then balances the dataset to only work with the same number of positive and negative samples with respect to binds\n    # selects only the molecule_smiles and binds columns\n    # writes the results as a parquet to write_data_dir\n    import dask.dataframe as dd\n    from dask.diagnostics import ProgressBar\n    raw_data = dd.read_parquet(raw_data_dir)\n    cols = ['molecule_smiles']\n    balance = False\n    if 'binds' in raw_data.columns:\n        balance = True\n        cols += ['binds']\n    filtered_data = raw_data[raw_data['protein_name'] == protein_name][cols]\n    if not balance:\n        with ProgressBar():\n            filtered_data.to_parquet(write_data_dir)\n            return None\n    positives = filtered_data[filtered_data['binds'] == 1]\n    with ProgressBar():\n        positives.to_parquet(write_data_dir)\n    positives = dd.read_parquet(write_data_dir)\n    negatives = filtered_data.query('binds == 0')\n    frac = positives.shape[0].compute() / len(negatives)\n    with ProgressBar():\n        negatives.sample(frac=frac).to_parquet(write_data_dir, append=True)\n\nprocess_data(TRAIN_RAW_DIR, 'train.parquet')\nprocess_data(TEST_RAW_DIR, 'test.parquet')","metadata":{"execution":{"iopub.status.busy":"2024-04-28T16:04:04.155395Z","iopub.execute_input":"2024-04-28T16:04:04.155817Z","iopub.status.idle":"2024-04-28T16:06:03.11052Z","shell.execute_reply.started":"2024-04-28T16:04:04.155783Z","shell.execute_reply":"2024-04-28T16:06:03.109478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from datasets import Dataset\ntrain = pd.read_parquet('train.parquet').sample(frac=1).reset_index()\ntrain.columns = ['id', 'text', 'labels']\ntest = pd.read_parquet('test.parquet').reset_index()\ntest.columns = ['id', 'text']\ntrain = Dataset.from_pandas(train)\ntest = Dataset.from_pandas(test)\ntrain_split = train.train_test_split(test_size=0.1)\ntrain =  train_split['train']\nval = train_split['test']","metadata":{"execution":{"iopub.status.busy":"2024-04-28T16:06:17.549012Z","iopub.execute_input":"2024-04-28T16:06:17.549423Z","iopub.status.idle":"2024-04-28T16:06:20.565602Z","shell.execute_reply.started":"2024-04-28T16:06:17.549391Z","shell.execute_reply":"2024-04-28T16:06:20.564592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoTokenizer\ntokenizer = AutoTokenizer.from_pretrained(BASE_MODEL)\ndef tokenize(input):\n    return tokenizer(input['text'], truncation=True, padding=True)\ntrain = train.map(tokenize, batched=True)\nval = val.map(tokenize, batched=True)\ntest = test.map(tokenize, batched=True)","metadata":{"execution":{"iopub.status.busy":"2024-04-28T16:06:27.816116Z","iopub.execute_input":"2024-04-28T16:06:27.817021Z","iopub.status.idle":"2024-04-28T16:11:07.369039Z","shell.execute_reply.started":"2024-04-28T16:06:27.816985Z","shell.execute_reply":"2024-04-28T16:11:07.367956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# load model and tokenizer required for finetuning\nfrom transformers import AutoModelForSequenceClassification\nmodel = AutoModelForSequenceClassification.from_pretrained(BASE_MODEL, num_labels=2)","metadata":{"execution":{"iopub.status.busy":"2024-04-28T16:11:17.470258Z","iopub.execute_input":"2024-04-28T16:11:17.471129Z","iopub.status.idle":"2024-04-28T16:11:19.886796Z","shell.execute_reply.started":"2024-04-28T16:11:17.471089Z","shell.execute_reply":"2024-04-28T16:11:19.885787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport evaluate\n\nmetric = evaluate.load(\"accuracy\")","metadata":{"execution":{"iopub.status.busy":"2024-04-28T16:11:19.888749Z","iopub.execute_input":"2024-04-28T16:11:19.889433Z","iopub.status.idle":"2024-04-28T16:11:32.368261Z","shell.execute_reply.started":"2024-04-28T16:11:19.889396Z","shell.execute_reply":"2024-04-28T16:11:32.367071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compute_metrics(eval_pred):\n    logits, labels = eval_pred\n    predictions = np.argmax(logits, axis=-1)\n    return metric.compute(predictions=predictions, references=labels)","metadata":{"execution":{"iopub.status.busy":"2024-04-28T16:11:32.370008Z","iopub.execute_input":"2024-04-28T16:11:32.370331Z","iopub.status.idle":"2024-04-28T16:11:32.380127Z","shell.execute_reply.started":"2024-04-28T16:11:32.370278Z","shell.execute_reply":"2024-04-28T16:11:32.379057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import TrainingArguments, Trainer\n\ntraining_args = TrainingArguments(output_dir=\"test_trainer\", evaluation_strategy=\"epoch\", report_to=None, per_device_train_batch_size=32)","metadata":{"execution":{"iopub.status.busy":"2024-04-28T16:11:38.953349Z","iopub.execute_input":"2024-04-28T16:11:38.954501Z","iopub.status.idle":"2024-04-28T16:11:39.052019Z","shell.execute_reply.started":"2024-04-28T16:11:38.95446Z","shell.execute_reply":"2024-04-28T16:11:39.051098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoTokenizer\ntokenizer = AutoTokenizer.from_pretrained(BASE_MODEL)\ntrainer = Trainer(\n    model=model,\n    args=training_args,\n    train_dataset=train,\n    eval_dataset=val,\n    compute_metrics=compute_metrics,\n    tokenizer=tokenizer\n)\ntrainer.train()","metadata":{"execution":{"iopub.status.busy":"2024-04-28T16:11:49.32684Z","iopub.execute_input":"2024-04-28T16:11:49.327256Z","iopub.status.idle":"2024-04-28T17:02:29.536817Z","shell.execute_reply.started":"2024-04-28T16:11:49.327225Z","shell.execute_reply":"2024-04-28T17:02:29.535803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TODO: calibrate model. for now just stick in your best model and see how well it performs\ntest_results = trainer.predict(test)","metadata":{"execution":{"iopub.status.busy":"2024-04-28T17:19:26.676769Z","iopub.execute_input":"2024-04-28T17:19:26.677194Z","iopub.status.idle":"2024-04-28T17:27:06.760056Z","shell.execute_reply.started":"2024-04-28T17:19:26.677157Z","shell.execute_reply":"2024-04-28T17:27:06.75907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def logit_to_prob(logits):\n    import math\n    return 1 / (1 + math.exp(-logits))\ntest_probs = [logit_to_prob(logit[1]) for logit in test_results.predictions]","metadata":{"execution":{"iopub.status.busy":"2024-04-28T17:29:43.583603Z","iopub.execute_input":"2024-04-28T17:29:43.58467Z","iopub.status.idle":"2024-04-28T17:29:44.268909Z","shell.execute_reply.started":"2024-04-28T17:29:43.584635Z","shell.execute_reply":"2024-04-28T17:29:44.26784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ids = pd.read_parquet('/kaggle/input/leash-BELKA/test.parquet')[['id']]","metadata":{"execution":{"iopub.status.busy":"2024-04-28T17:34:24.538386Z","iopub.execute_input":"2024-04-28T17:34:24.538753Z","iopub.status.idle":"2024-04-28T17:34:25.586232Z","shell.execute_reply.started":"2024-04-28T17:34:24.538726Z","shell.execute_reply":"2024-04-28T17:34:25.584988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ids['binds'] = test_probs","metadata":{"execution":{"iopub.status.busy":"2024-04-28T17:34:41.31693Z","iopub.execute_input":"2024-04-28T17:34:41.317877Z","iopub.status.idle":"2024-04-28T17:34:42.829382Z","shell.execute_reply.started":"2024-04-28T17:34:41.31784Z","shell.execute_reply":"2024-04-28T17:34:42.826779Z"},"trusted":true},"execution_count":null,"outputs":[]}]}