{"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":41875,"databundleVersionId":5521661,"sourceType":"competition"}],"dockerImageVersionId":30746,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import wandb\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nmy_secret = user_secrets.get_secret(\"weights_and_biases\")\nwandb.login(key=my_secret)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:50:34.606691Z","iopub.execute_input":"2024-07-23T04:50:34.607077Z","iopub.status.idle":"2024-07-23T04:50:37.477341Z","shell.execute_reply.started":"2024-07-23T04:50:34.607048Z","shell.execute_reply":"2024-07-23T04:50:37.476444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install biopython\n!pip install obonet\n!pip install evaluate","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:50:37.479319Z","iopub.execute_input":"2024-07-23T04:50:37.479977Z","iopub.status.idle":"2024-07-23T04:51:17.230423Z","shell.execute_reply.started":"2024-07-23T04:50:37.479944Z","shell.execute_reply":"2024-07-23T04:51:17.22926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\n\"\"\"for dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\"\"\"\nfrom Bio import SeqIO\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\nimport networkx\nimport obonet\nimport matplotlib.pyplot as plt\nimport seaborn as sns","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-07-23T04:51:17.232096Z","iopub.execute_input":"2024-07-23T04:51:17.232459Z","iopub.status.idle":"2024-07-23T04:51:18.606291Z","shell.execute_reply.started":"2024-07-23T04:51:17.232424Z","shell.execute_reply":"2024-07-23T04:51:18.605544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## EDA","metadata":{}},{"cell_type":"code","source":"def import_sequences(fasta_path):\n    records = list(SeqIO.parse(fasta_path, \"fasta\"))\n    data = []\n    for record in records:\n        data.append({\"id\": record.id, \"name\": record.name, \"description\": record.description, \"sequence\": str(record.seq)})\n    return pd.DataFrame(data)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:51:18.607541Z","iopub.execute_input":"2024-07-23T04:51:18.608191Z","iopub.status.idle":"2024-07-23T04:51:18.614022Z","shell.execute_reply.started":"2024-07-23T04:51:18.608157Z","shell.execute_reply":"2024-07-23T04:51:18.612943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sequences = import_sequences(\"/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta\")\nsequences","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:51:18.616494Z","iopub.execute_input":"2024-07-23T04:51:18.616769Z","iopub.status.idle":"2024-07-23T04:51:22.203173Z","shell.execute_reply.started":"2024-07-23T04:51:18.616748Z","shell.execute_reply":"2024-07-23T04:51:22.202263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"targets = pd.read_csv(\"/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv\", sep=\"\\t\")\ntargets","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:51:22.204538Z","iopub.execute_input":"2024-07-23T04:51:22.204893Z","iopub.status.idle":"2024-07-23T04:51:25.619318Z","shell.execute_reply.started":"2024-07-23T04:51:22.204861Z","shell.execute_reply":"2024-07-23T04:51:25.618407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"go_graph = obonet.read_obo(\"/kaggle/input/cafa-5-protein-function-prediction/Train/go-basic.obo\")\nnames = networkx.get_node_attributes(go_graph, \"name\")\ndf_terms = pd.json_normalize(names).T\ndf_terms[\"id\"] = df_terms.index\ndf_terms = df_terms.rename(columns={0:\"term\"})\ndf_terms = df_terms.reset_index(drop=True)[[\"id\", \"term\"]]\ndf_terms","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:51:25.620496Z","iopub.execute_input":"2024-07-23T04:51:25.620826Z","iopub.status.idle":"2024-07-23T04:51:44.261715Z","shell.execute_reply.started":"2024-07-23T04:51:25.620799Z","shell.execute_reply":"2024-07-23T04:51:44.260724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Lets explore some Molecular Functions","metadata":{}},{"cell_type":"code","source":"bpo = targets[targets.aspect == \"MFO\"]\ndata = bpo.merge(sequences[[\"id\", \"sequence\"]], left_on=\"EntryID\", right_on=\"id\")[[\"sequence\", \"term\"]]\ndata","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:51:44.262931Z","iopub.execute_input":"2024-07-23T04:51:44.26324Z","iopub.status.idle":"2024-07-23T04:51:45.566616Z","shell.execute_reply.started":"2024-07-23T04:51:44.263213Z","shell.execute_reply":"2024-07-23T04:51:45.56565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"count_terms = data.term.value_counts()\nfiltered_data = data\nsns.boxplot(count_terms)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:51:45.567968Z","iopub.execute_input":"2024-07-23T04:51:45.568304Z","iopub.status.idle":"2024-07-23T04:51:45.907266Z","shell.execute_reply.started":"2024-07-23T04:51:45.568277Z","shell.execute_reply":"2024-07-23T04:51:45.906431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"I will watch some proportions","metadata":{}},{"cell_type":"code","source":"filtered_data[\"values\"] = 1\nfiltered_data = filtered_data.drop_duplicates()\npivoted = filtered_data.pivot(columns=\"term\", index=\"sequence\", values=\"values\")\npivoted = pivoted.fillna(0)\nsorted_go = (pivoted.sum() / pivoted.count()).sort_values(ascending=False)\nsorted_go[0:50]","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:51:45.908511Z","iopub.execute_input":"2024-07-23T04:51:45.908849Z","iopub.status.idle":"2024-07-23T04:51:54.22885Z","shell.execute_reply.started":"2024-07-23T04:51:45.908821Z","shell.execute_reply":"2024-07-23T04:51:54.227899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_terms[df_terms.id.isin(sorted_go.index[0:20])]","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:51:54.23006Z","iopub.execute_input":"2024-07-23T04:51:54.230407Z","iopub.status.idle":"2024-07-23T04:51:54.251198Z","shell.execute_reply.started":"2024-07-23T04:51:54.230366Z","shell.execute_reply":"2024-07-23T04:51:54.250216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"I will focus on DNA binding","metadata":{}},{"cell_type":"code","source":"sorted_go.loc[\"GO:0003677\"]\npos = filtered_data[filtered_data.term == \"GO:0003677\"].drop_duplicates(subset=\"sequence\")\nneg = filtered_data[~filtered_data.sequence.isin(pos.sequence)].drop_duplicates(subset=\"sequence\").sample(len(pos))\nneg[\"values\"] = 0\ndna_binding = pd.concat([pos, neg])\ndna_binding = dna_binding.rename(columns={\"values\": \"dna\"}).drop(columns=[\"term\"]).drop_duplicates(subset=\"sequence\")\n#dna_binding = dna_binding.sample(1000)\ndna_binding[\"no-dna\"] = ~dna_binding.dna + 2\ndna_binding","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:51:54.252693Z","iopub.execute_input":"2024-07-23T04:51:54.252956Z","iopub.status.idle":"2024-07-23T04:51:54.533643Z","shell.execute_reply.started":"2024-07-23T04:51:54.252934Z","shell.execute_reply":"2024-07-23T04:51:54.53272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dna_binding[\"length\"] = dna_binding.sequence.str.len()\nq3 = dna_binding.length.quantile(0.75)\nq1 = dna_binding.length.quantile(0.25)\ndna_binding = dna_binding[dna_binding.length < q3 + 1.5*(q3 - q1)]\ndna_binding = dna_binding.drop(columns=[\"length\"])\ndna_binding","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:51:54.535003Z","iopub.execute_input":"2024-07-23T04:51:54.535398Z","iopub.status.idle":"2024-07-23T04:51:54.560264Z","shell.execute_reply.started":"2024-07-23T04:51:54.535347Z","shell.execute_reply":"2024-07-23T04:51:54.559422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sequences = dna_binding[\"sequence\"].tolist()\nlabels = dna_binding[\"dna\"].tolist()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:53:49.297854Z","iopub.execute_input":"2024-07-23T04:53:49.298851Z","iopub.status.idle":"2024-07-23T04:53:49.306548Z","shell.execute_reply.started":"2024-07-23T04:53:49.29881Z","shell.execute_reply":"2024-07-23T04:53:49.305277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Fine tune model","metadata":{}},{"cell_type":"code","source":"#Use huggingface model\nfrom transformers import AutoTokenizer, AutoModelForSequenceClassification, TrainingArguments, Trainer\n\nmodel_checkpoint=\"facebook/esm2_t6_8M_UR50D\"\n\ntokenizer = AutoTokenizer.from_pretrained(model_checkpoint)\nmodel = AutoModelForSequenceClassification.from_pretrained(model_checkpoint, num_labels=2)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:51:54.577158Z","iopub.execute_input":"2024-07-23T04:51:54.577449Z","iopub.status.idle":"2024-07-23T04:52:12.359747Z","shell.execute_reply.started":"2024-07-23T04:51:54.577418Z","shell.execute_reply":"2024-07-23T04:52:12.358865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:12.361161Z","iopub.execute_input":"2024-07-23T04:52:12.361893Z","iopub.status.idle":"2024-07-23T04:52:12.369422Z","shell.execute_reply.started":"2024-07-23T04:52:12.361864Z","shell.execute_reply":"2024-07-23T04:52:12.368189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\ntrain_sequences, test_sequences, train_labels, test_labels = train_test_split(sequences, labels, test_size=0.25, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:12.370611Z","iopub.execute_input":"2024-07-23T04:52:12.37148Z","iopub.status.idle":"2024-07-23T04:52:12.394716Z","shell.execute_reply.started":"2024-07-23T04:52:12.371449Z","shell.execute_reply":"2024-07-23T04:52:12.393881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_tokenized = tokenizer(train_sequences)\ntest_tokenized = tokenizer(test_sequences)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:12.395655Z","iopub.execute_input":"2024-07-23T04:52:12.395896Z","iopub.status.idle":"2024-07-23T04:52:34.058928Z","shell.execute_reply.started":"2024-07-23T04:52:12.395875Z","shell.execute_reply":"2024-07-23T04:52:34.058133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from datasets import Dataset\ntrain_dataset = Dataset.from_dict(train_tokenized)\ntest_dataset = Dataset.from_dict(test_tokenized)\ntrain_dataset = train_dataset.add_column(\"labels\", train_labels)\ntest_dataset = test_dataset.add_column(\"labels\", test_labels)\ntrain_dataset, test_dataset","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:34.060063Z","iopub.execute_input":"2024-07-23T04:52:34.060362Z","iopub.status.idle":"2024-07-23T04:52:36.376069Z","shell.execute_reply.started":"2024-07-23T04:52:34.060337Z","shell.execute_reply":"2024-07-23T04:52:36.375109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import TrainingArguments, Trainer\n\nmodel_name = model_checkpoint.split(\"/\")[-1]\nbatch_size = 8\nepochs = 5\nargs = TrainingArguments(\n    f\"{model_name}-finetuned-dna_binding\",\n    evaluation_strategy = \"epoch\",\n    save_strategy = \"epoch\",\n    learning_rate=1e-3,\n    per_device_train_batch_size=batch_size,\n    per_device_eval_batch_size=batch_size,\n    num_train_epochs=epochs,\n    weight_decay=0.01,\n    load_best_model_at_end=True,\n    metric_for_best_model=\"accuracy\",\n    do_train=True\n)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:36.37721Z","iopub.execute_input":"2024-07-23T04:52:36.377526Z","iopub.status.idle":"2024-07-23T04:52:36.496412Z","shell.execute_reply.started":"2024-07-23T04:52:36.377499Z","shell.execute_reply":"2024-07-23T04:52:36.49527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from evaluate import load\n\nmetric = load(\"accuracy\")\n\ndef compute_metrics(eval_pred):\n    predictions, labels = eval_pred\n    predictions = np.argmax(predictions, axis=1)\n    return metric.compute(predictions=predictions, references=labels)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:36.49744Z","iopub.execute_input":"2024-07-23T04:52:36.497744Z","iopub.status.idle":"2024-07-23T04:52:37.699291Z","shell.execute_reply.started":"2024-07-23T04:52:36.497719Z","shell.execute_reply":"2024-07-23T04:52:37.698405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for index, param in enumerate(model.parameters()):\n    param.requires_grad=False\n    \nfor index, param in enumerate(model.classifier.parameters()):\n    param.requires_grad=True","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:37.700428Z","iopub.execute_input":"2024-07-23T04:52:37.700715Z","iopub.status.idle":"2024-07-23T04:52:37.706403Z","shell.execute_reply.started":"2024-07-23T04:52:37.700691Z","shell.execute_reply":"2024-07-23T04:52:37.705132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pytorch_total_params = sum(p.numel() for p in model.parameters())\npytorch_total_params","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:37.707409Z","iopub.execute_input":"2024-07-23T04:52:37.707681Z","iopub.status.idle":"2024-07-23T04:52:37.718581Z","shell.execute_reply.started":"2024-07-23T04:52:37.707659Z","shell.execute_reply":"2024-07-23T04:52:37.717656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pytorch_non_trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad == False)\npytorch_non_trainable_params","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:37.71948Z","iopub.execute_input":"2024-07-23T04:52:37.719767Z","iopub.status.idle":"2024-07-23T04:52:37.730113Z","shell.execute_reply.started":"2024-07-23T04:52:37.719745Z","shell.execute_reply":"2024-07-23T04:52:37.729238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pytorch_trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\npytorch_trainable_params","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:37.731308Z","iopub.execute_input":"2024-07-23T04:52:37.73165Z","iopub.status.idle":"2024-07-23T04:52:37.743473Z","shell.execute_reply.started":"2024-07-23T04:52:37.731621Z","shell.execute_reply":"2024-07-23T04:52:37.742583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pytorch_non_trainable_params + pytorch_trainable_params == pytorch_total_params","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:37.744365Z","iopub.execute_input":"2024-07-23T04:52:37.744655Z","iopub.status.idle":"2024-07-23T04:52:37.752066Z","shell.execute_reply.started":"2024-07-23T04:52:37.744633Z","shell.execute_reply":"2024-07-23T04:52:37.751206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Trainer(\n    model,\n    args,\n    train_dataset=train_dataset,\n    eval_dataset=test_dataset,\n    tokenizer=tokenizer,\n    compute_metrics=compute_metrics,\n)","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:37.753169Z","iopub.execute_input":"2024-07-23T04:52:37.753518Z","iopub.status.idle":"2024-07-23T04:52:37.945744Z","shell.execute_reply.started":"2024-07-23T04:52:37.753488Z","shell.execute_reply":"2024-07-23T04:52:37.944894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wandb.init()\ntrainer.train()\n","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:37.946877Z","iopub.execute_input":"2024-07-23T04:52:37.947571Z","iopub.status.idle":"2024-07-23T04:52:58.841934Z","shell.execute_reply.started":"2024-07-23T04:52:37.947501Z","shell.execute_reply":"2024-07-23T04:52:58.840447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport gc\n#del model\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-07-23T04:52:58.842982Z","iopub.status.idle":"2024-07-23T04:52:58.84336Z","shell.execute_reply.started":"2024-07-23T04:52:58.843176Z","shell.execute_reply":"2024-07-23T04:52:58.843193Z"},"trusted":true},"execution_count":null,"outputs":[]}]}