{"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":9260546,"sourceType":"datasetVersion","datasetId":5580487}],"dockerImageVersionId":30684,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport time\nt0start = time.time() \n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport os","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-08-25T16:00:30.894304Z","iopub.execute_input":"2024-08-25T16:00:30.894992Z","iopub.status.idle":"2024-08-25T16:00:33.059028Z","shell.execute_reply.started":"2024-08-25T16:00:30.894961Z","shell.execute_reply":"2024-08-25T16:00:33.058137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pickle","metadata":{"execution":{"iopub.status.busy":"2024-08-25T16:00:33.070991Z","iopub.execute_input":"2024-08-25T16:00:33.071775Z","iopub.status.idle":"2024-08-25T16:00:33.076055Z","shell.execute_reply.started":"2024-08-25T16:00:33.071745Z","shell.execute_reply":"2024-08-25T16:00:33.074933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load SMILES data","metadata":{}},{"cell_type":"code","source":"%%time\ntrain = pd.read_csv('/kaggle/input/predictddi-dataset/filtered_processed_full.csv')\n#train = df[(df['smiles'].notnull()) & (df['formula'].notnull())]\ntrain.info()\n","metadata":{"execution":{"iopub.status.busy":"2024-08-25T16:00:35.946333Z","iopub.execute_input":"2024-08-25T16:00:35.947022Z","iopub.status.idle":"2024-08-25T16:00:35.990984Z","shell.execute_reply.started":"2024-08-25T16:00:35.946988Z","shell.execute_reply":"2024-08-25T16:00:35.989933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nimport torch\nfrom transformers import AutoModelForMaskedLM, AutoTokenizer\n\nchemberta = AutoModelForMaskedLM.from_pretrained(\"DeepChem/ChemBERTa-77M-MTR\")\ntokenizer = AutoTokenizer.from_pretrained(\"DeepChem/ChemBERTa-77M-MTR\")\n","metadata":{"execution":{"iopub.status.busy":"2024-08-25T16:01:00.619932Z","iopub.execute_input":"2024-08-25T16:01:00.620662Z","iopub.status.idle":"2024-08-25T16:01:13.130778Z","shell.execute_reply.started":"2024-08-25T16:01:00.620632Z","shell.execute_reply":"2024-08-25T16:01:13.129774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"embedding.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_output[0].shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nfrom tqdm import tqdm\n\nchemberta.eval()\ndef featurize_ChemBERTa(smiles_list, padding=True):\n    embeddings_cls = torch.zeros(len(smiles_list), 600)\n\n    with torch.no_grad():\n        for i, smiles in enumerate(tqdm(smiles_list)):\n            if smiles is not None:\n                encoded_input = tokenizer(smiles, return_tensors=\"pt\",padding=padding,truncation=True)\n                model_output = chemberta(**encoded_input)\n\n                embedding = model_output[0][::,0,::]\n                embeddings_cls[i] = embedding\n            \n    return embeddings_cls.numpy()\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create embeddings for SMILES data\n%%time\nN = int(1e4)\nsmiles_list = train['smiles'].iloc[:N].to_list()\nprint(smiles_list[:10])\ntrain_cls_pad_true = featurize_ChemBERTa(smiles_list)\n\n# CPU:  int(1e4) -  Wall time: 1min 47s\n# GPU:  int(1e4) -  Wall time: 1min 16s\n","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_cls_pad_true.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n# Save embeddings\nnp.save('smiles_embedd.npy', train_cls_pad_true)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# False sampling","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch\n\n\nold_train_df = pd.read_csv('/kaggle/input/predictddi-dataset/train.csv')\nfiltered_df = pd.read_csv('/kaggle/input/predictddi-dataset/filtered_processed_full.csv')\nsmiles_embeddings = np.load('/kaggle/input/predictddi-dataset/smiles_embedd.npy')","metadata":{"execution":{"iopub.status.busy":"2024-08-27T16:50:10.604960Z","iopub.execute_input":"2024-08-27T16:50:10.605334Z","iopub.status.idle":"2024-08-27T16:50:10.635239Z","shell.execute_reply.started":"2024-08-27T16:50:10.605305Z","shell.execute_reply":"2024-08-27T16:50:10.634195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(smiles_embeddings.shape) ","metadata":{"execution":{"iopub.status.busy":"2024-08-27T16:50:13.121654Z","iopub.execute_input":"2024-08-27T16:50:13.122038Z","iopub.status.idle":"2024-08-27T16:50:13.126981Z","shell.execute_reply.started":"2024-08-27T16:50:13.122009Z","shell.execute_reply":"2024-08-27T16:50:13.126007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nfrom sklearn.decomposition import PCA","metadata":{"execution":{"iopub.status.busy":"2024-08-27T16:50:14.953996Z","iopub.execute_input":"2024-08-27T16:50:14.954851Z","iopub.status.idle":"2024-08-27T16:50:16.257472Z","shell.execute_reply.started":"2024-08-27T16:50:14.954812Z","shell.execute_reply":"2024-08-27T16:50:16.256554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Reduce the dimensionality of embeddings\npca = PCA(n_components=128)\nsmiles_embeddings_reduced = pca.fit_transform(smiles_embeddings)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T16:50:23.313776Z","iopub.execute_input":"2024-08-27T16:50:23.314551Z","iopub.status.idle":"2024-08-27T16:50:23.744741Z","shell.execute_reply.started":"2024-08-27T16:50:23.314516Z","shell.execute_reply":"2024-08-27T16:50:23.743447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(smiles_embeddings_reduced.shape)  # Should print (2090, 64)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T16:50:26.411550Z","iopub.execute_input":"2024-08-27T16:50:26.412161Z","iopub.status.idle":"2024-08-27T16:50:26.416476Z","shell.execute_reply.started":"2024-08-27T16:50:26.412130Z","shell.execute_reply":"2024-08-27T16:50:26.415522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# training drug pairs\ntrain_pairs = list()\nfor _, row in old_train_df.iterrows():\n    train_pairs.append((row['name1'], row['name2']))\n\n# all drug pairs\nall_drugs = filtered_df['name'].to_list()\nall_pairs = list()\nall_drugs_length = len(all_drugs)\nfor i in range(all_drugs_length - 1):\n    for j in range(i + 1, all_drugs_length):  # Start j at i+1 to avoid duplicate pairs like (A, B) and (B, A)\n        pair = (all_drugs[i], all_drugs[j])\n        reverse_pair = (all_drugs[j], all_drugs[i])\n        \n        # Check if neither the pair nor its reverse is in train_pairs\n        if pair not in train_pairs and reverse_pair not in train_pairs:\n            all_pairs.append(pair)\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T16:56:45.843759Z","iopub.execute_input":"2024-08-27T16:56:45.844371Z","iopub.status.idle":"2024-08-27T17:01:03.949706Z","shell.execute_reply.started":"2024-08-27T16:56:45.844338Z","shell.execute_reply":"2024-08-27T17:01:03.948699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('all_pairs len: ', len(all_pairs))\nprint('train_pairs len: ', len(train_pairs))","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:01:03.951380Z","iopub.execute_input":"2024-08-27T17:01:03.951670Z","iopub.status.idle":"2024-08-27T17:01:03.956418Z","shell.execute_reply.started":"2024-08-27T17:01:03.951645Z","shell.execute_reply":"2024-08-27T17:01:03.955590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dict from name to embedding\nname_to_embedding = {}\n\nfor index, row in filtered_df.iterrows():\n    name = row['name']\n    embedding = smiles_embeddings_reduced[index]\n    name_to_embedding[name] = embedding","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:03:49.733251Z","iopub.execute_input":"2024-08-27T17:03:49.734239Z","iopub.status.idle":"2024-08-27T17:03:49.851546Z","shell.execute_reply.started":"2024-08-27T17:03:49.734202Z","shell.execute_reply":"2024-08-27T17:03:49.850556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# get relation embedding from a pair\ndef get_relation_embedding(pair):\n    name1, name2 = pair\n    embedding1 = name_to_embedding.get(name1, np.zeros(64))  \n    embedding2 = name_to_embedding.get(name2, np.zeros(64)) \n    relation_embedding = np.concatenate([embedding1, embedding2])\n    return relation_embedding","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:03:53.332745Z","iopub.execute_input":"2024-08-27T17:03:53.333384Z","iopub.status.idle":"2024-08-27T17:03:53.338501Z","shell.execute_reply.started":"2024-08-27T17:03:53.333352Z","shell.execute_reply":"2024-08-27T17:03:53.337572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dict from pair to embedding\npair_to_embedding = {}\nfor pair in all_pairs:\n    embedding = get_relation_embedding(pair)\n    pair_to_embedding[pair] = embedding","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:09:31.084995Z","iopub.execute_input":"2024-08-27T17:09:31.085716Z","iopub.status.idle":"2024-08-27T17:09:41.608291Z","shell.execute_reply.started":"2024-08-27T17:09:31.085683Z","shell.execute_reply":"2024-08-27T17:09:41.607292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Reverse dictionary to map embeddings back to pairs\nembedding_to_pair = {tuple(v): k for k, v in pair_to_embedding.items()}\n\ndef get_pair_from_embedding(embedding):\n    return embedding_to_pair.get(tuple(embedding), (\"Unknown\", \"Unknown\"))","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:11:15.030412Z","iopub.execute_input":"2024-08-27T17:11:15.030669Z","iopub.status.idle":"2024-08-27T17:12:54.244723Z","shell.execute_reply.started":"2024-08-27T17:11:15.030646Z","shell.execute_reply":"2024-08-27T17:12:54.243950Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# group training dataframe by labels\ngrouped = old_train_df.groupby('label')","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:14:07.474675Z","iopub.execute_input":"2024-08-27T17:14:07.475367Z","iopub.status.idle":"2024-08-27T17:14:07.479880Z","shell.execute_reply.started":"2024-08-27T17:14:07.475337Z","shell.execute_reply":"2024-08-27T17:14:07.478861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# label to centeroid\nlabel_to_centroid = {}\n\nfor label, group in grouped:\n    embeddings = np.array([get_relation_embedding((row['name1'], row['name2'])) for _, row in group.iterrows()])\n    centroid = np.mean(embeddings, axis=0)\n    label_to_centroid[label] = centroid","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:14:14.901707Z","iopub.execute_input":"2024-08-27T17:14:14.902348Z","iopub.status.idle":"2024-08-27T17:14:15.080595Z","shell.execute_reply.started":"2024-08-27T17:14:14.902318Z","shell.execute_reply":"2024-08-27T17:14:15.079846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from scipy.spatial.distance import cdist\n\n# Get centroids for all three labels\ncentroid_1 = label_to_centroid['mechanism']  \ncentroid_2 = label_to_centroid['effect']\ncentroid_3= label_to_centroid['advise']\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:17:20.983933Z","iopub.execute_input":"2024-08-27T17:17:20.984258Z","iopub.status.idle":"2024-08-27T17:17:20.989458Z","shell.execute_reply.started":"2024-08-27T17:17:20.984233Z","shell.execute_reply":"2024-08-27T17:17:20.988330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"min_distance = max(np.linalg.norm(centroid_2 - centroid_1), np.linalg.norm(centroid_3 - centroid_1), np.linalg.norm(centroid_2 - centroid_3))","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:17:24.323595Z","iopub.execute_input":"2024-08-27T17:17:24.324064Z","iopub.status.idle":"2024-08-27T17:17:24.329291Z","shell.execute_reply.started":"2024-08-27T17:17:24.324032Z","shell.execute_reply":"2024-08-27T17:17:24.328347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"min_distance","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:17:41.691601Z","iopub.execute_input":"2024-08-27T17:17:41.692210Z","iopub.status.idle":"2024-08-27T17:17:41.698741Z","shell.execute_reply.started":"2024-08-27T17:17:41.692177Z","shell.execute_reply":"2024-08-27T17:17:41.697762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get valid pairs\nall_distances = []\n\nfor pair, embedding in pair_to_embedding.items():\n    # Calculate distance from all centroids\n    dist_1 = np.linalg.norm(embedding - centroid_1) \n    dist_2 = np.linalg.norm(embedding - centroid_2)\n    dist_3 = np.linalg.norm(embedding - centroid_3)\n    \n    if (dist_1 >= min_distance and dist_2 >= min_distance and dist_3 >= min_distance):\n        # Append to list\n        all_distances.append(pair)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:23:12.945600Z","iopub.execute_input":"2024-08-27T17:23:12.946416Z","iopub.status.idle":"2024-08-27T17:23:52.276817Z","shell.execute_reply.started":"2024-08-27T17:23:12.946384Z","shell.execute_reply":"2024-08-27T17:23:52.276003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(all_distances)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:24:55.533727Z","iopub.execute_input":"2024-08-27T17:24:55.534583Z","iopub.status.idle":"2024-08-27T17:24:55.539979Z","shell.execute_reply.started":"2024-08-27T17:24:55.534553Z","shell.execute_reply":"2024-08-27T17:24:55.539086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_false_samples = 16200","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:26:14.223924Z","iopub.execute_input":"2024-08-27T17:26:14.224633Z","iopub.status.idle":"2024-08-27T17:26:14.228522Z","shell.execute_reply.started":"2024-08-27T17:26:14.224601Z","shell.execute_reply":"2024-08-27T17:26:14.227611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\n# randomly sampling from the list\nfalse_pairs = random.sample(all_distances, num_false_samples)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:27:22.843705Z","iopub.execute_input":"2024-08-27T17:27:22.844075Z","iopub.status.idle":"2024-08-27T17:27:22.867385Z","shell.execute_reply.started":"2024-08-27T17:27:22.844045Z","shell.execute_reply":"2024-08-27T17:27:22.866554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(false_pairs)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:27:39.953927Z","iopub.execute_input":"2024-08-27T17:27:39.954275Z","iopub.status.idle":"2024-08-27T17:27:39.960085Z","shell.execute_reply.started":"2024-08-27T17:27:39.954245Z","shell.execute_reply":"2024-08-27T17:27:39.959180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"non_df = pd.DataFrame(false_pairs, columns=['name1', 'name2'])\n\n# Add a 'label' column with all values set to 'non'\nnon_df['label'] = 'non'\n\n# Reorder columns to match the desired output\nnon_df = non_df[['label', 'name1', 'name2']]","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:34:28.856981Z","iopub.execute_input":"2024-08-27T17:34:28.857340Z","iopub.status.idle":"2024-08-27T17:34:28.869929Z","shell.execute_reply.started":"2024-08-27T17:34:28.857310Z","shell.execute_reply":"2024-08-27T17:34:28.869030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.to_csv('train.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-08-26T17:43:07.270448Z","iopub.execute_input":"2024-08-26T17:43:07.270832Z","iopub.status.idle":"2024-08-26T17:43:07.288377Z","shell.execute_reply.started":"2024-08-26T17:43:07.270802Z","shell.execute_reply":"2024-08-26T17:43:07.287330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Classify","metadata":{}},{"cell_type":"code","source":"#Importing all libraries\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom torch import optim\n\nfrom torchvision import datasets, transforms\nfrom torch.utils.data import DataLoader, TensorDataset, Dataset\nfrom torchvision.utils import make_grid\nfrom torch.autograd import Variable\n\nimport time\nimport helper\nfrom tqdm import tqdm\n\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:34:34.094742Z","iopub.execute_input":"2024-08-27T17:34:34.095598Z","iopub.status.idle":"2024-08-27T17:34:36.978954Z","shell.execute_reply.started":"2024-08-27T17:34:34.095567Z","shell.execute_reply":"2024-08-27T17:34:36.978168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dir = '/kaggle/input/predictddi-dataset/new_train.csv'\ntest_dir = '/kaggle/input/predictddi-dataset/new_test.csv'","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:34:38.633783Z","iopub.execute_input":"2024-08-27T17:34:38.634316Z","iopub.status.idle":"2024-08-27T17:34:38.638690Z","shell.execute_reply.started":"2024-08-27T17:34:38.634286Z","shell.execute_reply":"2024-08-27T17:34:38.637758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.concat([old_train_df, non_df], ignore_index=True)\ntest_df = pd.read_csv(test_dir)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:35:55.623781Z","iopub.execute_input":"2024-08-27T17:35:55.624710Z","iopub.status.idle":"2024-08-27T17:35:55.640134Z","shell.execute_reply.started":"2024-08-27T17:35:55.624663Z","shell.execute_reply":"2024-08-27T17:35:55.639384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_df)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:07.661973Z","iopub.execute_input":"2024-08-27T17:36:07.662342Z","iopub.status.idle":"2024-08-27T17:36:07.668271Z","shell.execute_reply.started":"2024-08-27T17:36:07.662312Z","shell.execute_reply":"2024-08-27T17:36:07.667347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(test_df)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:15.293821Z","iopub.execute_input":"2024-08-27T17:36:15.294180Z","iopub.status.idle":"2024-08-27T17:36:15.300523Z","shell.execute_reply.started":"2024-08-27T17:36:15.294151Z","shell.execute_reply":"2024-08-27T17:36:15.299599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n# Split the DataFrame into training and validation sets\ntrain_df, val_df = train_test_split(train_df, test_size=0.1, random_state=42)\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:24.283576Z","iopub.execute_input":"2024-08-27T17:36:24.283949Z","iopub.status.idle":"2024-08-27T17:36:24.294928Z","shell.execute_reply.started":"2024-08-27T17:36:24.283918Z","shell.execute_reply":"2024-08-27T17:36:24.294034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Print the number of rows in each set\nprint(f\"Number of rows in train_df: {len(train_df)}\")\nprint(f\"Number of rows in val_df: {len(val_df)}\")\nprint((f\"Number of rows in test_df: {len(test_df)}\"))","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:26.063590Z","iopub.execute_input":"2024-08-27T17:36:26.063961Z","iopub.status.idle":"2024-08-27T17:36:26.069196Z","shell.execute_reply.started":"2024-08-27T17:36:26.063932Z","shell.execute_reply":"2024-08-27T17:36:26.068278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\n\nclass DrugDataset(Dataset):\n    def __init__(self, df):\n        self.df = df\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        # Extract the label and the pair (name1, name2) from the DataFrame\n        row = self.df.iloc[idx]\n        label_map = {\n            'mechanism': 0,\n            'effect': 1,\n            'advise': 2,\n            'non': 3\n        }\n        label = label_map[row['label']]\n        pair = (row['name1'], row['name2'])\n        \n        # Get the embedding of the relation using the provided function\n        embedding_np = get_relation_embedding(pair)\n        \n        # Convert the NumPy array to a PyTorch tensor\n        embedding = torch.tensor(embedding_np, dtype=torch.float)\n        \n        return embedding, label\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:29.234686Z","iopub.execute_input":"2024-08-27T17:36:29.235057Z","iopub.status.idle":"2024-08-27T17:36:29.242739Z","shell.execute_reply.started":"2024-08-27T17:36:29.235029Z","shell.execute_reply":"2024-08-27T17:36:29.241815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = DrugDataset(train_df)\nval_dataset = DrugDataset(val_df)\ntest_dataset = DrugDataset(test_df)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:32.833679Z","iopub.execute_input":"2024-08-27T17:36:32.834051Z","iopub.status.idle":"2024-08-27T17:36:32.838572Z","shell.execute_reply.started":"2024-08-27T17:36:32.834020Z","shell.execute_reply":"2024-08-27T17:36:32.837595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(train_dataset))","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:35.043626Z","iopub.execute_input":"2024-08-27T17:36:35.043998Z","iopub.status.idle":"2024-08-27T17:36:35.048697Z","shell.execute_reply.started":"2024-08-27T17:36:35.043968Z","shell.execute_reply":"2024-08-27T17:36:35.047748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_sizes = {'train': len(train_dataset), 'val': len(val_dataset), 'test': len(test_dataset)}","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:37.153524Z","iopub.execute_input":"2024-08-27T17:36:37.153884Z","iopub.status.idle":"2024-08-27T17:36:37.158187Z","shell.execute_reply.started":"2024-08-27T17:36:37.153855Z","shell.execute_reply":"2024-08-27T17:36:37.157317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_sizes['train']","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:38.713688Z","iopub.execute_input":"2024-08-27T17:36:38.714053Z","iopub.status.idle":"2024-08-27T17:36:38.719709Z","shell.execute_reply.started":"2024-08-27T17:36:38.714024Z","shell.execute_reply":"2024-08-27T17:36:38.718874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Dataloader\ntrain_loader = DataLoader(train_dataset, batch_size = 64, num_workers=2, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size = 64, num_workers=2, shuffle=True)\ntest_loader = DataLoader(test_dataset, batch_size = 64, num_workers=2, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:41.413865Z","iopub.execute_input":"2024-08-27T17:36:41.414222Z","iopub.status.idle":"2024-08-27T17:36:41.419671Z","shell.execute_reply.started":"2024-08-27T17:36:41.414195Z","shell.execute_reply":"2024-08-27T17:36:41.418729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataloader = {'train':train_loader, 'val':val_loader, 'test': test_loader}","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:45.833664Z","iopub.execute_input":"2024-08-27T17:36:45.834513Z","iopub.status.idle":"2024-08-27T17:36:45.838439Z","shell.execute_reply.started":"2024-08-27T17:36:45.834480Z","shell.execute_reply":"2024-08-27T17:36:45.837445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_inputs = 256\nn_outputs = 4","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:51.153591Z","iopub.execute_input":"2024-08-27T17:36:51.154059Z","iopub.status.idle":"2024-08-27T17:36:51.159561Z","shell.execute_reply.started":"2024-08-27T17:36:51.154027Z","shell.execute_reply":"2024-08-27T17:36:51.158730Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Multi-layer perceptron\nmodel = nn.Sequential(nn.Linear(n_inputs, 512),\n                      nn.LeakyReLU(),\n                      nn.Linear(512, 256),\n                      nn.BatchNorm1d(256),\n                      nn.LeakyReLU(),\n                      nn.Linear(256, n_outputs),\n                      nn.LogSoftmax(dim=1))\n\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:54.514689Z","iopub.execute_input":"2024-08-27T17:36:54.515064Z","iopub.status.idle":"2024-08-27T17:36:54.566697Z","shell.execute_reply.started":"2024-08-27T17:36:54.515034Z","shell.execute_reply":"2024-08-27T17:36:54.565963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\" )\ndevice","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:56.973859Z","iopub.execute_input":"2024-08-27T17:36:56.974220Z","iopub.status.idle":"2024-08-27T17:36:57.019162Z","shell.execute_reply.started":"2024-08-27T17:36:56.974189Z","shell.execute_reply":"2024-08-27T17:36:57.018248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.to(device)\nmodel","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:36:59.133505Z","iopub.execute_input":"2024-08-27T17:36:59.133878Z","iopub.status.idle":"2024-08-27T17:36:59.287428Z","shell.execute_reply.started":"2024-08-27T17:36:59.133848Z","shell.execute_reply":"2024-08-27T17:36:59.286538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lr = 1e-5","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:37:01.363791Z","iopub.execute_input":"2024-08-27T17:37:01.364199Z","iopub.status.idle":"2024-08-27T17:37:01.368499Z","shell.execute_reply.started":"2024-08-27T17:37:01.364160Z","shell.execute_reply":"2024-08-27T17:37:01.367534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=lr)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:37:03.573846Z","iopub.execute_input":"2024-08-27T17:37:03.574202Z","iopub.status.idle":"2024-08-27T17:37:03.579422Z","shell.execute_reply.started":"2024-08-27T17:37:03.574174Z","shell.execute_reply":"2024-08-27T17:37:03.578427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training function\ndef train_model(model, criterion, optimizer, num_epochs, save_model_path, dataloader):\n\n    for epoch in range(num_epochs):\n        print(f'Epoch {epoch}/{num_epochs - 1}')\n        print('-' * 10)\n\n        # Each epoch has a training and validation phase\n        for phase in ['train', 'val']:\n            if phase == 'train':\n                model.train()  # Set model to training mode\n            else:\n                model.eval()   # Set model to evaluate mode\n\n            running_loss = 0.0\n            running_corrects = 0\n\n            # Iterate over data.\n            for inputs, labels in dataloader[phase]:\n                inputs = inputs.to(device)\n                labels = labels.to(device)\n\n                # zero the parameter gradients\n                optimizer.zero_grad()\n\n                # forward\n                # track history if only in train\n                with torch.set_grad_enabled(phase == 'train'):\n                    outputs = model(inputs)\n                    _, preds = torch.max(outputs, 1)\n                    loss = criterion(outputs, labels)\n\n                    # backward + optimize only if in training phase\n                    if phase == 'train':\n                        loss.backward()\n                        optimizer.step()\n\n                # statistics\n                running_loss += loss.item() * inputs.size(0)\n                running_corrects += torch.sum(preds == labels.data)\n#             if phase == 'train':\n#                 scheduler.step()\n\n            epoch_loss = running_loss / dataset_sizes[phase]\n            epoch_acc = running_corrects.double() / dataset_sizes[phase]\n\n            print(f'{phase} Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}')\n            phase_acc, phase_loss = phase + \"acc\", phase + \"loss\"\n#             wandb.log({str(phase_acc): epoch_acc, str(phase_loss):epoch_loss})\n        torch.save(model.state_dict(), save_model_path)\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:37:05.363452Z","iopub.execute_input":"2024-08-27T17:37:05.363830Z","iopub.status.idle":"2024-08-27T17:37:05.374523Z","shell.execute_reply.started":"2024-08-27T17:37:05.363779Z","shell.execute_reply":"2024-08-27T17:37:05.373559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 200\nsave_model_path = 'drug_1.pt'","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:37:12.344972Z","iopub.execute_input":"2024-08-27T17:37:12.345681Z","iopub.status.idle":"2024-08-27T17:37:12.349568Z","shell.execute_reply.started":"2024-08-27T17:37:12.345651Z","shell.execute_reply":"2024-08-27T17:37:12.348606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_model(model, criterion, optimizer, num_epochs, save_model_path, dataloader)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:37:14.493614Z","iopub.execute_input":"2024-08-27T17:37:14.494206Z","iopub.status.idle":"2024-08-27T17:45:34.148619Z","shell.execute_reply.started":"2024-08-27T17:37:14.494173Z","shell.execute_reply":"2024-08-27T17:45:34.147478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import precision_score, recall_score, f1_score","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:45:37.595565Z","iopub.execute_input":"2024-08-27T17:45:37.596521Z","iopub.status.idle":"2024-08-27T17:45:37.600976Z","shell.execute_reply.started":"2024-08-27T17:45:37.596472Z","shell.execute_reply":"2024-08-27T17:45:37.600034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_model(model, dataloader):\n    # Set the model to evaluation mode\n    model.eval()\n\n    # Initialize lists to store true labels and predictions\n    all_labels = []\n    all_preds = []\n    \n    # A list to store the results\n    results = []\n\n    for inputs, labels in dataloader['test']:\n        inputs, labels = inputs.to(device), labels.to(device)\n\n        # Get the predicted labels\n        outputs = model(inputs)\n        _, preds = torch.max(outputs, 1)\n\n        # Store the true labels and predictions\n        all_labels.extend(labels.cpu().numpy())\n        all_preds.extend(preds.cpu().numpy())\n\n    # Calculate precision, recall, and F1-score\n    precision = precision_score(all_labels, all_preds, average='weighted')\n    recall = recall_score(all_labels, all_preds, average='weighted')\n    f1 = f1_score(all_labels, all_preds, average='weighted')\n\n    # Print the results\n    print(f'Precision: {precision:.4f}')\n    print(f'Recall:    {recall:.4f}')\n    print(f'F1 Score:  {f1:.4f}')","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:45:40.864927Z","iopub.execute_input":"2024-08-27T17:45:40.865543Z","iopub.status.idle":"2024-08-27T17:45:40.872964Z","shell.execute_reply.started":"2024-08-27T17:45:40.865509Z","shell.execute_reply":"2024-08-27T17:45:40.872088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_model(model, dataloader)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T17:45:44.954343Z","iopub.execute_input":"2024-08-27T17:45:44.954990Z","iopub.status.idle":"2024-08-27T17:45:45.629369Z","shell.execute_reply.started":"2024-08-27T17:45:44.954954Z","shell.execute_reply":"2024-08-27T17:45:45.628267Z"},"trusted":true},"execution_count":null,"outputs":[]}]}