{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport torch.nn as nn\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport plotly.express as px","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-20T01:20:39.210557Z","iopub.execute_input":"2023-06-20T01:20:39.210942Z","iopub.status.idle":"2023-06-20T01:20:39.216642Z","shell.execute_reply.started":"2023-06-20T01:20:39.210911Z","shell.execute_reply":"2023-06-20T01:20:39.215435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:39.227339Z","iopub.execute_input":"2023-06-20T01:20:39.228215Z","iopub.status.idle":"2023-06-20T01:20:39.23253Z","shell.execute_reply.started":"2023-06-20T01:20:39.228182Z","shell.execute_reply":"2023-06-20T01:20:39.231552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv('/kaggle/input/c5-dataset/train_data.csv')\n# discard rows with long sequences\ntrain_df['seq_len'] = train_df.seq.map(lambda s: len(s))\nMAX_SEQ_LEN = 1000\ntrain_df = train_df[train_df.seq_len <= MAX_SEQ_LEN].reset_index(drop=True)\ntrain_df","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:39.25017Z","iopub.execute_input":"2023-06-20T01:20:39.250423Z","iopub.status.idle":"2023-06-20T01:20:41.79352Z","shell.execute_reply.started":"2023-06-20T01:20:39.250401Z","shell.execute_reply":"2023-06-20T01:20:41.792628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_amios = set(ch for s in train_df.seq for ch in s)\nidx_to_amino = list(all_amios)\nPAD_IDX = len(idx_to_amino)\nidx_to_amino.append('<pad>')\namino_to_idx = {v: k for k,v in enumerate(idx_to_amino)}\nlen(idx_to_amino)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:41.795336Z","iopub.execute_input":"2023-06-20T01:20:41.796862Z","iopub.status.idle":"2023-06-20T01:20:41.80678Z","shell.execute_reply.started":"2023-06-20T01:20:41.796826Z","shell.execute_reply":"2023-06-20T01:20:41.805576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"amino_tensors = [torch.tensor([amino_to_idx[ch] for ch in s], device=DEVICE) for s in train_df.seq]\namino_tensors[0]","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:41.808306Z","iopub.execute_input":"2023-06-20T01:20:41.808679Z","iopub.status.idle":"2023-06-20T01:20:44.460172Z","shell.execute_reply.started":"2023-06-20T01:20:41.808645Z","shell.execute_reply":"2023-06-20T01:20:44.459081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"term_strings = [ts[2:-2].split(\"', '\") for ts in train_df.term]\nall_terms = set(term for ts in term_strings for term in ts)\nidx_to_term = list(all_terms)\nterm_to_idx = {v: k for k,v in enumerate(idx_to_term)}\nlen(idx_to_term)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:44.46296Z","iopub.execute_input":"2023-06-20T01:20:44.463323Z","iopub.status.idle":"2023-06-20T01:20:44.472596Z","shell.execute_reply.started":"2023-06-20T01:20:44.463291Z","shell.execute_reply":"2023-06-20T01:20:44.471538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"term_tensors = [torch.tensor([term_to_idx[term] for term in ts], device=DEVICE) for ts in term_strings]\nterm_tensors[0]","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:44.474349Z","iopub.execute_input":"2023-06-20T01:20:44.474695Z","iopub.status.idle":"2023-06-20T01:20:44.493359Z","shell.execute_reply.started":"2023-06-20T01:20:44.474665Z","shell.execute_reply":"2023-06-20T01:20:44.492559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"onehot = torch.eye(len(idx_to_term), device=DEVICE)\nclass C5Dataset(torch.utils.data.Dataset):\n    def __init__(self, indicies):\n        self.indicies = indicies\n    def __len__(self):\n        return len(self.indicies)\n    def __getitem__(self, i):\n        i = self.indicies[i]\n        target = onehot[term_tensors[i]].sum(0)\n        return amino_tensors[i], target\n\nC5Dataset(indicies=torch.arange(len(term_tensors), device=DEVICE))[0]","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:44.494545Z","iopub.execute_input":"2023-06-20T01:20:44.494856Z","iopub.status.idle":"2023-06-20T01:20:44.55942Z","shell.execute_reply.started":"2023-06-20T01:20:44.494828Z","shell.execute_reply":"2023-06-20T01:20:44.558575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def collate_fn(batch):\n    lens = torch.tensor([len(b[0]) for b in batch], device=DEVICE)\n    \n    # get max sequence length\n    max_length = lens.max().item()\n    \n    # pad each element to max length\n    inputs = [\n        torch.concat([\n            e[0], \n            torch.ones(max_length - len(e[0]), device=DEVICE, dtype=torch.long) * PAD_IDX,\n        ]) for e in batch\n    ]\n    return (torch.stack(inputs), lens), torch.stack([e[1] for e in batch])\n\n\nexample_batch = next(iter(torch.utils.data.DataLoader(\n    C5Dataset(indicies=torch.arange(len(term_tensors), device=DEVICE)), \n    batch_size=2, \n    collate_fn=collate_fn\n)))\nexample_batch[0][0].shape, example_batch[0][1].shape, example_batch[1].shape","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:44.560834Z","iopub.execute_input":"2023-06-20T01:20:44.561162Z","iopub.status.idle":"2023-06-20T01:20:44.583257Z","shell.execute_reply.started":"2023-06-20T01:20:44.561133Z","shell.execute_reply":"2023-06-20T01:20:44.582426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataloader = torch.utils.data.DataLoader(\n    C5Dataset(indicies=torch.arange(len(term_tensors), device=DEVICE)), \n    batch_size=32, \n    collate_fn=collate_fn\n)\n\nfor _ in tqdm(range(1)):\n    for i, (inputs, targets) in enumerate(test_dataloader):\n        pass","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:44.584569Z","iopub.execute_input":"2023-06-20T01:20:44.584881Z","iopub.status.idle":"2023-06-20T01:20:44.609949Z","shell.execute_reply.started":"2023-06-20T01:20:44.584853Z","shell.execute_reply":"2023-06-20T01:20:44.608982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"x = torch.tensor([\n    [0, 1, 2],\n    [3, 4, 5]\n])\n\nonehot_x = torch.eye(3).bool()\n\nx[onehot_x[[1, 2]]]","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:44.61137Z","iopub.execute_input":"2023-06-20T01:20:44.611702Z","iopub.status.idle":"2023-06-20T01:20:44.625691Z","shell.execute_reply.started":"2023-06-20T01:20:44.611672Z","shell.execute_reply":"2023-06-20T01:20:44.624872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# amino embedding + LSTM + linear projection to terms\nclass C5Model(nn.Module):\n    def __init__(self, embedding_dim, hidden_dim, lstm_num_layers):\n        super().__init__()\n        self.embedding = nn.Embedding(num_embeddings=len(idx_to_amino), embedding_dim=embedding_dim)\n        self.lstm = nn.LSTM(input_size=embedding_dim, hidden_size=hidden_dim, num_layers=lstm_num_layers, batch_first=True)\n        self.linear_out = nn.Linear(in_features=hidden_dim, out_features=len(idx_to_term))\n        self.register_buffer('onehot', torch.eye(MAX_SEQ_LEN).bool(), persistent=False)\n\n    def forward(self, x):\n        x, lens = x\n        x = self.embedding(x)\n        x = self.lstm(x)[0]\n        x = x[self.onehot[lens - 1][:, :lens.max()]] # get the output at the end of each sequence (ignore padding)\n        return self.linear_out(torch.relu(x))\n\n\nexample_model = C5Model(embedding_dim=64, hidden_dim=64, lstm_num_layers=1).to(DEVICE)\nexample_model(example_batch[0]).shape # 2xlen(idx_to_term)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:44.62989Z","iopub.execute_input":"2023-06-20T01:20:44.630688Z","iopub.status.idle":"2023-06-20T01:20:46.759489Z","shell.execute_reply.started":"2023-06-20T01:20:44.630657Z","shell.execute_reply":"2023-06-20T01:20:46.758536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"param_counts = [layer.numel() for layer in example_model.parameters()]\nparam_counts, sum(param_counts)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:46.760777Z","iopub.execute_input":"2023-06-20T01:20:46.761202Z","iopub.status.idle":"2023-06-20T01:20:46.768109Z","shell.execute_reply.started":"2023-06-20T01:20:46.761169Z","shell.execute_reply":"2023-06-20T01:20:46.767125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"def train_val_split(batch_size):\n    n = len(term_tensors)\n    indices = torch.randperm(n, device=DEVICE)\n    n_train = int(0.8 * n)\n    \n    train_dataloader = torch.utils.data.DataLoader(\n        C5Dataset(indicies=indices[0:n_train]), \n        batch_size=batch_size, \n        collate_fn=collate_fn,\n        shuffle=True,\n    )\n    \n    val_dataloader = torch.utils.data.DataLoader(\n        C5Dataset(indicies=indices[n_train:]), \n        batch_size=batch_size, \n        collate_fn=collate_fn,\n        shuffle=False,\n    )\n    \n    return train_dataloader, val_dataloader","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:46.769559Z","iopub.execute_input":"2023-06-20T01:20:46.770211Z","iopub.status.idle":"2023-06-20T01:20:46.780586Z","shell.execute_reply.started":"2023-06-20T01:20:46.770179Z","shell.execute_reply":"2023-06-20T01:20:46.779594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader, val_dataloader = train_val_split(32)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:46.782656Z","iopub.execute_input":"2023-06-20T01:20:46.78291Z","iopub.status.idle":"2023-06-20T01:20:46.792938Z","shell.execute_reply.started":"2023-06-20T01:20:46.782888Z","shell.execute_reply":"2023-06-20T01:20:46.792008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = C5Model(embedding_dim=64, hidden_dim=256, lstm_num_layers=2).to(DEVICE)\nlossfn = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\n\nNUM_EPOCHS = 100\ntrain_losses = []\nval_losses = []\n\nfor epoch in tqdm(range(NUM_EPOCHS)):\n    model.train()\n    epoch_train_losses = []\n    for inputs, targets in train_dataloader:\n        logits = model(inputs)\n        loss = lossfn(logits, targets)\n        loss.backward()\n        optimizer.step()\n        optimizer.zero_grad()\n        epoch_train_losses.append(loss.detach())\n    train_losses.extend(torch.stack(epoch_train_losses).to('cpu').numpy())\n    \n    model.eval()\n    epoch_val_losses = []\n    with torch.no_grad():\n        for inputs, targets in val_dataloader:\n            logits = model(inputs)\n            loss = lossfn(logits, targets)\n            epoch_val_losses.append(loss)\n    val_losses.append(torch.stack(epoch_val_losses).mean().item())\n    \n    if (epoch + 1) % 10 == 0:\n        torch.save(model.state_dict(), f'model_weights_epoch_{epoch}.pt')","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:20:46.794205Z","iopub.execute_input":"2023-06-20T01:20:46.794582Z","iopub.status.idle":"2023-06-20T01:21:18.99669Z","shell.execute_reply.started":"2023-06-20T01:20:46.794552Z","shell.execute_reply":"2023-06-20T01:21:18.99576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"px.line(train_losses)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:21:18.998167Z","iopub.execute_input":"2023-06-20T01:21:18.998765Z","iopub.status.idle":"2023-06-20T01:21:20.288895Z","shell.execute_reply.started":"2023-06-20T01:21:18.998727Z","shell.execute_reply":"2023-06-20T01:21:20.288001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"px.line(val_losses)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:21:20.290339Z","iopub.execute_input":"2023-06-20T01:21:20.290713Z","iopub.status.idle":"2023-06-20T01:21:20.361767Z","shell.execute_reply.started":"2023-06-20T01:21:20.290681Z","shell.execute_reply":"2023-06-20T01:21:20.360761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), 'model_weights.pt')","metadata":{"execution":{"iopub.status.busy":"2023-06-20T01:21:20.363023Z","iopub.execute_input":"2023-06-20T01:21:20.363423Z","iopub.status.idle":"2023-06-20T01:21:20.37984Z","shell.execute_reply.started":"2023-06-20T01:21:20.363392Z","shell.execute_reply":"2023-06-20T01:21:20.379018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Infer","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}