{"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":8635058,"sourceType":"datasetVersion","datasetId":5143792}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Basic SMILES vs. AtomInSmiles (AIS) Tokenization","metadata":{"execution":{"iopub.status.busy":"2024-06-10T16:15:21.571653Z","iopub.execute_input":"2024-06-10T16:15:21.573055Z","iopub.status.idle":"2024-06-10T16:15:21.608862Z","shell.execute_reply.started":"2024-06-10T16:15:21.573011Z","shell.execute_reply":"2024-06-10T16:15:21.606948Z"}}},{"cell_type":"markdown","source":"When molecular data in the form of SMILES strings is fed to a sequence model, a typical way of tokenizing the string is to create one token for each type of atom in the molecule (see for example the [deepchem SmilesTokenizer](https://deepchem.readthedocs.io/en/2.4.0/api_reference/tokenizers.html)). However, that way of tokenization results in just a few embedding vectors per type of atom\\* making it difficult for the model to account for the fact that an atom's properties are strongly influenced by its chemical environment within the molecule. [Ucak et *al.*](https://doi.org/10.1186/s13321-023-00725-9) recently introduced [AtomInSmiles](https://github.com/snu-lcbc/atom-in-SMILES), a more expressive tokenization scheme that creates a greater diversity of tokens taking into account all neighboring atoms of a given atom during tokenization. \n\nIn this notebook I use a small subset of the [LeashBio BELKA](https://www.kaggle.com/competitions/leash-BELKA) dataset (1.5M examples) to show that using AIS tokenization instead of SMILES tokenization leads to better results for classification of molecule binding across all three proteins using a 1D convolutional neural network with skip connections. My experiments show that the average precision score can be significantly improved by 14-50 % using AtomInSmiles istead of basic SMILES tokenization. The superiority of AIS tokenization likely extrapolates to using greater amounts of data, other sequence model architectures or longer training times.\n\nThe [dataset](https://www.kaggle.com/datasets/frenio/leashbio-belka-numericalized-smiles-and-ais) used contains numericalized and zero-padded arrays for the complete [LeashBio BELKA](https://www.kaggle.com/competitions/leash-BELKA) dataset using both SMILES and AIS tokenization schemes. It is adapted from @shlomoron's [BELKA: Shrunken train set](https://www.kaggle.com/datasets/shlomoron/belka-shrunken-train-set) who thankfully increased the usability of the original data.\n\nThe [training framework (atai)](https://github.com/frenio/atai) used is adapted from the [miniai framework](https://github.com/fastai/course22p2/tree/master) from the [fast.ai 2022](https://course.fast.ai/Lessons/part2.html) course.\n\n\\*Note that it does differentiate between aromatic and non-aromatic, charge states, and a few other exceptions.","metadata":{}},{"cell_type":"markdown","source":"## 1. Installation of Dependencies and Imports","metadata":{}},{"cell_type":"code","source":"!pip -qq install atai\n!pip -qq install rdkit\n!pip -qq install atomInSmiles","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:37:10.003073Z","iopub.execute_input":"2024-06-14T14:37:10.003437Z","iopub.status.idle":"2024-06-14T14:37:51.132158Z","shell.execute_reply.started":"2024-06-14T14:37:10.003404Z","shell.execute_reply":"2024-06-14T14:37:51.130881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from functools import partial\n\nimport torch\nimport torch.nn.functional as F\nimport torch.nn as nn\n\n\nfrom torcheval.metrics import BinaryAccuracy, Mean, BinaryAUROC\nfrom torcheval.metrics.functional import binary_auroc, binary_accuracy\nfrom torchmetrics.classification import BinaryMatthewsCorrCoef, BinaryAveragePrecision\nfrom torchmetrics.functional.classification import matthews_corrcoef, binary_average_precision\n\nimport numpy as np\nimport pandas as pd\nimport random\nimport atomInSmiles\nimport fastcore.all as fc\nfrom rdkit import Chem\nfrom rdkit.Chem import Draw\nfrom atai.core import *","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:37:51.134275Z","iopub.execute_input":"2024-06-14T14:37:51.134592Z","iopub.status.idle":"2024-06-14T14:38:00.723196Z","shell.execute_reply.started":"2024-06-14T14:37:51.134561Z","shell.execute_reply":"2024-06-14T14:38:00.722333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Functions","metadata":{}},{"cell_type":"markdown","source":"Defining a function to set random seeds to ensure deterministic training runs for better comparibility of results.","metadata":{}},{"cell_type":"code","source":"def random_seed(seed_value, use_cuda):\n    np.random.seed(seed_value) # for numpy random\n    torch.manual_seed(seed_value) # for pytorch\n    random.seed(seed_value) # for python random\n    if use_cuda:\n        torch.cuda.manual_seed(seed_value)\n        torch.cuda.manual_seed_all(seed_value)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n        \nseed_value=42\nrandom_seed(seed_value, use_cuda=True)\nprint(f\"Random seed test: {torch.randint(0, 1000, (1,)).item()}\")","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:38:00.724788Z","iopub.execute_input":"2024-06-14T14:38:00.725409Z","iopub.status.idle":"2024-06-14T14:38:00.755225Z","shell.execute_reply.started":"2024-06-14T14:38:00.725371Z","shell.execute_reply":"2024-06-14T14:38:00.754267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The [LeashBio BELKA](https://www.kaggle.com/competitions/leash-BELKA) dataset is very imbalanced containing only about 0.5% of positive (binding) examples. The `get_train_valid_data` function is designed to handle this imbalance and prepare the data for training and validation.\n\nThe function takes four arguments: the number of negative examples to be used, the name of the label column for a protein (`'binds_BRD4'`, `'binds_HSA'` or `'binds_sEH'`), the name of a tokenization technique (`'smiles'` or `'ais'`), and a random seed. It retrieves the numericalized training data using the specified number of negative examples.\n\nThe function returns validation data with a positive ratio that is representative of the complete dataset. All other positive examples, along with the specified number of negative examples, are returned as training data. It is important to note that the training data is returned unshuffled and needs to be shuffled during the creation of dataloaders to ensure effective training.","metadata":{}},{"cell_type":"code","source":"## bind column names: 'binds_BRD4', 'binds_HSA', 'binds_sEH'\n## tokenization types: 'smiles', 'ais'\ndef get_train_valid_data(neg_size=1_000_000, bind_column='binds_BRD4', tok_type=\"smiles\", random_seed=seed_value):\n    print(f\"Loading numericalized {tok_type.upper()} data...\")\n    trn = np.load(f\"/kaggle/input/leashbio-belka-numericalized-smiles-and-ais/train_numericalized_{tok_type}.npz\")['arr_0']\n    print(f\"All features shape: {trn.shape}\")\n    print(f\"Loading bind labels for protein {bind_column[6:]}...\")\n    lbs = pd.read_parquet(f\"/kaggle/input/leashbio-belka-numericalized-smiles-and-ais/train_data_smiles_and_ais.parquet\", columns=[bind_column])[bind_column].to_numpy()\n    print(f\"All labels shape: {lbs.shape}\")\n    mask = lbs == 1\n    neg = trn[~mask]\n    pos = trn[mask]\n    print(\"Cleaning up!\")\n    import gc\n    del trn      \n    gc.collect()\n    neg_size=int(neg_size)\n    rng = torch.Generator()\n    rng.manual_seed(random_seed)\n    neg_idx = torch.randperm(neg.shape[0], generator=rng)\n    pos_idx = torch.randperm(pos.shape[0], generator=rng)\n    pos_rate = lbs.sum()/neg.shape[0]\n    print(f\"Positive-negative ratio: {pos_rate:.4f}\")\n    neg_size_valid = neg_size // 10\n    pos_size_valid = int(neg_size_valid * pos_rate)\n    print(f\"Extracting a train set with {neg_size} negative examples and {pos.shape[0] - pos_size_valid} positive examples...\")\n    xs = torch.tensor(np.concatenate([neg[neg_idx[:neg_size]], pos[pos_idx[:-pos_size_valid]]]))\n    ys = torch.tensor([0] * neg_size + [1] * (pos.shape[0] - pos_size_valid), dtype=torch.float)\n    print(f\"Features shape: {xs.shape}\\nLabels sum: {ys.sum()}\")\n    print(f\"Extracting a valid set with {neg_size_valid} negative examples and {pos_size_valid} positive examples...\")\n    valid_xs = torch.tensor(np.concatenate([neg[neg_idx[neg_size:neg_size+neg_size_valid]], pos[pos_idx[-pos_size_valid:]]]))\n    valid_ys = torch.tensor([0] * neg_size_valid + [1] * pos_size_valid, dtype=torch.float)\n    print(f\"Valid features shape: {valid_xs.shape}\\nValid labels sum: {valid_ys.sum()}\")\n    print(\"Cleaning up!\")\n    del neg\n    del pos\n    gc.collect()\n    print(f\"Extracting vocab size...\")\n    vocab = open(f\"/kaggle/input/leashbio-belka-numericalized-smiles-and-ais/{tok_type}_vocab.txt\", 'r').read().splitlines()\n    vocab_size = len(vocab)\n    print(f\"Vocab size: {vocab_size}\")\n    print(\"Done!\")\n    return xs, ys, valid_xs, valid_ys, vocab_size","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:38:00.757469Z","iopub.execute_input":"2024-06-14T14:38:00.757814Z","iopub.status.idle":"2024-06-14T14:38:00.771961Z","shell.execute_reply.started":"2024-06-14T14:38:00.757787Z","shell.execute_reply":"2024-06-14T14:38:00.770894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The code below defines 1D convolution layers with skip connections. It's based on the 2D ResBlock from [Lesson 18](https://course.fast.ai/Lessons/lesson18.html) of the 2022 fast.ai course, but adapted for 1D data.","metadata":{}},{"cell_type":"code","source":"def conv1d(ni, nf, ks=3, stride=2, act=nn.ReLU, norm=None, bias=None):\n    if bias is None: bias = not isinstance(norm, (nn.BatchNorm1d,nn.BatchNorm2d,nn.BatchNorm3d))\n    layers = [nn.Conv1d(ni, nf, stride=stride, kernel_size=ks, padding=ks//2, bias=bias)]\n    if norm: layers.append(norm(nf))\n    if act: layers.append(act())\n    return nn.Sequential(*layers)\n\ndef _conv1d_block(ni, nf, stride, act=nn.ReLU, norm=None, ks=3):\n    return nn.Sequential(conv1d(ni, nf, stride=1, act=act, norm=norm, ks=ks),\n                         conv1d(nf, nf, stride=stride, act=None, norm=norm, ks=ks))\n\nclass ResBlock1d(nn.Module):\n    def __init__(self, ni, nf, stride=1, ks=3, act=nn.ReLU, norm=None):\n        super().__init__()\n        self.convs = _conv1d_block(ni, nf, stride=stride, ks=ks, act=act, norm=norm)\n        self.idconv = fc.noop if ni==nf else conv1d(ni, nf, stride=1, ks=1, act=None)\n        self.pool = fc.noop if stride==1 else nn.AvgPool1d(stride, ceil_mode=True)\n        self.act = act()\n\n    def forward(self, x): return self.act(self.convs(x) + self.pool(self.idconv(x)))","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:38:00.773239Z","iopub.execute_input":"2024-06-14T14:38:00.773602Z","iopub.status.idle":"2024-06-14T14:38:00.788745Z","shell.execute_reply.started":"2024-06-14T14:38:00.773572Z","shell.execute_reply":"2024-06-14T14:38:00.787783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The `Reshape` layer is needed to switch tensor ranks into the correct order for the following layers.","metadata":{}},{"cell_type":"code","source":"class Reshape(nn.Module):\n    def forward(self, x):\n        B, L, C = x.shape\n        return x.view(B, C, L)","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:38:00.789844Z","iopub.execute_input":"2024-06-14T14:38:00.790173Z","iopub.status.idle":"2024-06-14T14:38:00.804551Z","shell.execute_reply.started":"2024-06-14T14:38:00.790146Z","shell.execute_reply":"2024-06-14T14:38:00.803480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"A leaky ReLU with a negative offset is used to improve training.","metadata":{}},{"cell_type":"code","source":"act_genrelu = partial(GeneralRelu, leak=0.1, sub=0.4)","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:38:00.805635Z","iopub.execute_input":"2024-06-14T14:38:00.805933Z","iopub.status.idle":"2024-06-14T14:38:00.814961Z","shell.execute_reply.started":"2024-06-14T14:38:00.805910Z","shell.execute_reply":"2024-06-14T14:38:00.814151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The `get_model` function is used to quickly reinitialize the model.","metadata":{}},{"cell_type":"code","source":"def get_model(vocab_size, n_embd, dropout):\n    model = nn.Sequential(nn.Embedding(vocab_size, n_embd, padding_idx=0), Reshape(),\n                      ResBlock1d(n_embd, 8, ks=15, stride=2, norm=nn.BatchNorm1d, act=act_genrelu), nn.Dropout(dropout),\n                      ResBlock1d(8, 16, ks=13, stride=2, norm=nn.BatchNorm1d, act=act_genrelu), nn.Dropout(dropout),\n                      ResBlock1d(16, 32, ks=11, stride=2, norm=nn.BatchNorm1d, act=act_genrelu), nn.Dropout(dropout),\n                      ResBlock1d(32, 32, ks=9, stride=2, norm=nn.BatchNorm1d, act=act_genrelu), nn.Dropout(dropout),\n                      ResBlock1d(32, 64, ks=7, stride=2, norm=nn.BatchNorm1d, act=act_genrelu), nn.Dropout(dropout),\n                      ResBlock1d(64, 64, ks=5, stride=2, norm=nn.BatchNorm1d, act=act_genrelu), nn.Dropout(dropout),\n                      ResBlock1d(64, 128, ks=3, stride=2, norm=nn.BatchNorm1d, act=act_genrelu), nn.Dropout(dropout),\n                      nn.Flatten(1, -1),\n                      nn.Linear(128, 1),\n                      nn.Flatten(0, -1))\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:38:00.816174Z","iopub.execute_input":"2024-06-14T14:38:00.816470Z","iopub.status.idle":"2024-06-14T14:38:00.826201Z","shell.execute_reply.started":"2024-06-14T14:38:00.816445Z","shell.execute_reply":"2024-06-14T14:38:00.825296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The code below contains small adjustments to deal with data type issues.","metadata":{}},{"cell_type":"code","source":"class BinaryAP(BinaryAveragePrecision):\n    def update(self, preds, target) -> None:\n        # Convert target tensor to integer dtype\n        target = target.to(dtype=torch.long)\n        # Call the original update method with the converted target tensor\n        super().update(preds, target)\n        \nclass MetricsWithLogitsCB(MetricsCB):\n    def after_batch(self, learn):\n        x, y = to_cpu(learn.batch)\n        for m in self.metrics.values(): m.update(to_cpu(F.sigmoid(learn.preds)), y)\n        self.loss.update(to_cpu(learn.loss), weight=len(x))\n        \nclass BELKATrainLearner(TrainLearner):\n    def predict(self): self.preds = self.model(self.batch[0].to(torch.long))","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:38:00.827292Z","iopub.execute_input":"2024-06-14T14:38:00.827604Z","iopub.status.idle":"2024-06-14T14:38:00.841452Z","shell.execute_reply.started":"2024-06-14T14:38:00.827579Z","shell.execute_reply":"2024-06-14T14:38:00.840586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Take a Look at the Data","metadata":{}},{"cell_type":"markdown","source":"This section is meant to show that everything went well during numericalization and the numericalized molecules can be correctly decoded into molecular structures.","metadata":{}},{"cell_type":"code","source":"indices = [random.randint(0, 98_000_000) for i in range(9)]\nsmiles = np.load('/kaggle/input/leashbio-belka-numericalized-smiles-and-ais/train_numericalized_smiles.npz')['arr_0'][indices]\nais = np.load('/kaggle/input/leashbio-belka-numericalized-smiles-and-ais/train_numericalized_ais.npz')['arr_0'][indices]","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:38:00.844932Z","iopub.execute_input":"2024-06-14T14:38:00.845237Z","iopub.status.idle":"2024-06-14T14:39:15.207292Z","shell.execute_reply.started":"2024-06-14T14:38:00.845210Z","shell.execute_reply":"2024-06-14T14:39:15.206323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 3.1 SMILES Decoded Structures","metadata":{}},{"cell_type":"markdown","source":"The variable `smiles` contains nine numericalized molecules randomly sampled from the train data.","metadata":{}},{"cell_type":"code","source":"smiles[0], smiles.shape","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:39:15.208567Z","iopub.execute_input":"2024-06-14T14:39:15.208864Z","iopub.status.idle":"2024-06-14T14:39:15.217302Z","shell.execute_reply.started":"2024-06-14T14:39:15.208839Z","shell.execute_reply":"2024-06-14T14:39:15.216268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The code below uses the `smiles_vocab.txt` to decode the numericalized molecules into SMILES.","metadata":{}},{"cell_type":"code","source":"vocab = open('/kaggle/input/leashbio-belka-numericalized-smiles-and-ais/smiles_vocab.txt', 'r').read().splitlines()\nstr2int = {st:i for i, st in enumerate(vocab)}\nint2str = {i:st for i, st in enumerate(vocab)}\nencode = lambda s: np.array([str2int[st] for st in s], dtype=np.uint8)\ndecode = lambda l: (''.join([int2str[i] for i in l]))\n\nsmiles_decoded = list(map(decode, smiles))\nsmiles_decoded[:3]","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:39:15.218442Z","iopub.execute_input":"2024-06-14T14:39:15.218754Z","iopub.status.idle":"2024-06-14T14:39:15.240056Z","shell.execute_reply.started":"2024-06-14T14:39:15.218718Z","shell.execute_reply":"2024-06-14T14:39:15.238997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The structures can be displayed using `rdkit.Chem`.","metadata":{}},{"cell_type":"code","source":"Draw.MolsToGridImage([Chem.MolFromSmiles(x) for x in smiles_decoded], molsPerRow=3, subImgSize=(400,300))","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:39:15.241591Z","iopub.execute_input":"2024-06-14T14:39:15.242110Z","iopub.status.idle":"2024-06-14T14:39:15.351607Z","shell.execute_reply.started":"2024-06-14T14:39:15.242075Z","shell.execute_reply":"2024-06-14T14:39:15.350579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 3.2 AIS Decoded Structures","metadata":{}},{"cell_type":"markdown","source":"The variable `ais` contains the same nine numericalized molecules randomly drawn from the train data as the `smiles` variable, but the numbers are different, because of the different tokenization scheme used.","metadata":{}},{"cell_type":"code","source":"ais[0], ais.shape","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:39:15.353416Z","iopub.execute_input":"2024-06-14T14:39:15.353792Z","iopub.status.idle":"2024-06-14T14:39:15.360795Z","shell.execute_reply.started":"2024-06-14T14:39:15.353760Z","shell.execute_reply":"2024-06-14T14:39:15.359914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The numericalized data is decoded to AIS strings using the vocab in `ais_vocab.txt` and the code below.","metadata":{}},{"cell_type":"code","source":"vocab = open('/kaggle/input/leashbio-belka-numericalized-smiles-and-ais/ais_vocab.txt', 'r').read().splitlines()\nstr2int = {st:i for i, st in enumerate(vocab)}\nint2str = {i:st for i, st in enumerate(vocab)}\nencode = lambda s: np.array([str2int[st] for st in s], dtype=np.uint8)\ndecode = lambda l: (' '.join([int2str[i] for i in l]))\n\nais_decoded = list(map(decode, ais))\nais_decoded[:3]","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:39:15.361944Z","iopub.execute_input":"2024-06-14T14:39:15.362216Z","iopub.status.idle":"2024-06-14T14:39:15.380651Z","shell.execute_reply.started":"2024-06-14T14:39:15.362193Z","shell.execute_reply":"2024-06-14T14:39:15.379754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The AIS strings can be translated back into SMILES using the decode function in the [atomInSmiles](https://github.com/snu-lcbc/atom-in-SMILES) package. Note that the decoder doesn't add the brackets around dysprosium (`Dy` instead of `[Dy]`) which is used as a placeholder for the DNA tag in the BELKA dataset. This leads to issues when displaying the structures using `rdkit`, so the brackets need to be added manually.","metadata":{}},{"cell_type":"code","source":"ais_to_smiles = list(map(atomInSmiles.decode, ais_decoded))\nais_to_smiles = list(map(lambda x: x.replace('Dy', '[Dy]'), ais_to_smiles))\nais_to_smiles[:3]","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:39:15.381805Z","iopub.execute_input":"2024-06-14T14:39:15.382145Z","iopub.status.idle":"2024-06-14T14:39:15.390812Z","shell.execute_reply.started":"2024-06-14T14:39:15.382111Z","shell.execute_reply":"2024-06-14T14:39:15.389813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now the structures can be displayed using `rdkit`.","metadata":{}},{"cell_type":"code","source":"Draw.MolsToGridImage([Chem.MolFromSmiles(x) for x in ais_to_smiles], molsPerRow=3, subImgSize=(400,300))","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:39:15.392005Z","iopub.execute_input":"2024-06-14T14:39:15.392317Z","iopub.status.idle":"2024-06-14T14:39:15.498558Z","shell.execute_reply.started":"2024-06-14T14:39:15.392272Z","shell.execute_reply":"2024-06-14T14:39:15.496594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Visual inspection of the random sample of structures shows that decoding the numericalized data leads to the same molecular structures for both tokenization methods, so everything seems to be in order.","metadata":{}},{"cell_type":"markdown","source":"## 4. Basic SMILES vs. AIS Comparison","metadata":{}},{"cell_type":"markdown","source":"In the following section 6 models are trained – one for each combination of tokenization method (SMILES and AIS) and proteins in the [LeashBio BELKA](https://www.kaggle.com/competitions/leash-BELKA) dataset (BRD4, HSA, and sEH).","metadata":{}},{"cell_type":"markdown","source":"### 4.1 BRD4 Binding Classification","metadata":{}},{"cell_type":"markdown","source":"The models are all trained for 5 epochs using a learning rate of 0.001, an embedding size of 32, 20% dropout, and a batch size of 256.","metadata":{}},{"cell_type":"code","source":"epochs = 5\nlr = 1e-3\nn_embd = 32\ndropout = 0.2\nbs = 256","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:39:15.500389Z","iopub.execute_input":"2024-06-14T14:39:15.500809Z","iopub.status.idle":"2024-06-14T14:39:15.507828Z","shell.execute_reply.started":"2024-06-14T14:39:15.500771Z","shell.execute_reply":"2024-06-14T14:39:15.506420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### 4.1.1 Basic SMILES Tokenization","metadata":{}},{"cell_type":"markdown","source":"The code below resets the random seed and tests if the randomness has been properly reset. Note that `torch.randint(0, 1000, (1,))` picks the same number as in Section 2.","metadata":{}},{"cell_type":"code","source":"random_seed(seed_value, use_cuda=True)\nprint(f\"Random seed test: {torch.randint(0, 1000, (1,)).item()}\")","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:39:15.509367Z","iopub.execute_input":"2024-06-14T14:39:15.510013Z","iopub.status.idle":"2024-06-14T14:39:15.519096Z","shell.execute_reply.started":"2024-06-14T14:39:15.509975Z","shell.execute_reply":"2024-06-14T14:39:15.517860Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Next, the function `get_train_valid_data` is used to load train data for protein binding data for the protein BRD4 using basic SMILES tokenization using a small subset of 1,000,000 negative examples for training. Features and labels are combined into `Dataset`s and the positive weight (`pos_weight`) to be applied in the loss function is calculated from the ratio of negative to positive examples in the train data.","metadata":{}},{"cell_type":"code","source":"xs, ys, valid_xs, valid_ys, vocab_size = get_train_valid_data(neg_size=1e6,\n                                                              bind_column='binds_BRD4',\n                                                              tok_type='smiles')\ntrnds = Dataset(xs, ys)\nvldds = Dataset(valid_xs, valid_ys)\npos_weight = (ys.shape[0] - ys.sum()) / ys.sum()\nprint(f\"Positive weight: {pos_weight:.2f}\")","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:39:15.520237Z","iopub.execute_input":"2024-06-14T14:39:15.520519Z","iopub.status.idle":"2024-06-14T14:40:01.081475Z","shell.execute_reply.started":"2024-06-14T14:39:15.520496Z","shell.execute_reply":"2024-06-14T14:40:01.080530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The function `get_dls` creates the dataloaders and shuffles the train data. Kaiming He initialization is applied to the initial model weights (`iw`). The loss_function takes `pos_weight` into account. Calbacks (`cbs`) are defined to send data to GPU, show training progress, and calculate metrics. Finally, the learner is initialized, the sum of parameters of the model is printed, and the model is trained.","metadata":{}},{"cell_type":"code","source":"dls = get_dls(trnds, vldds, bs=bs)\n\nmodel = get_model(vocab_size, n_embd, dropout)\niw = partial(init_weights, leaky=0.1)\nmodel = model.apply(iw)\n\nloss_func = partial(F.binary_cross_entropy_with_logits, pos_weight=pos_weight)\nmetrics = MetricsWithLogitsCB(BinaryAccuracy(), BinaryMatthewsCorrCoef(), BinaryAP(), BinaryAUROC())\ncbs = [DeviceCB(), ProgressCB(plot=False), metrics]\n\nlearn = BELKATrainLearner(model, dls, loss_func, lr=lr, cbs=cbs, opt_func=torch.optim.AdamW)\nprint(f\"Parameters total: {sum(p.nelement() for p in model.parameters())}\")\nlearn.fit(epochs)","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:40:01.082769Z","iopub.execute_input":"2024-06-14T14:40:01.083073Z","iopub.status.idle":"2024-06-14T14:58:00.999565Z","shell.execute_reply.started":"2024-06-14T14:40:01.083046Z","shell.execute_reply":"2024-06-14T14:58:00.998651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### 4.1.2 AIS Tokenization","metadata":{}},{"cell_type":"markdown","source":"The same procedure as in 4.1.1, but using AIS tokenization.","metadata":{}},{"cell_type":"code","source":"random_seed(seed_value, use_cuda=True)\nprint(f\"Random seed test: {torch.randint(0, 1000, (1,)).item()}\")","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:58:01.000844Z","iopub.execute_input":"2024-06-14T14:58:01.001141Z","iopub.status.idle":"2024-06-14T14:58:01.006719Z","shell.execute_reply.started":"2024-06-14T14:58:01.001113Z","shell.execute_reply":"2024-06-14T14:58:01.005812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Prepare AIS data and calculate the positive weight.","metadata":{}},{"cell_type":"code","source":"xs, ys, valid_xs, valid_ys, vocab_size = get_train_valid_data(neg_size=1e6,\n                                                              bind_column='binds_BRD4',\n                                                              tok_type='ais')\ntrnds = Dataset(xs, ys)\nvldds = Dataset(valid_xs, valid_ys)\npos_weight = (ys.shape[0] - ys.sum()) / ys.sum()\nprint(f\"Positive weight: {pos_weight:.2f}\")","metadata":{"execution":{"iopub.status.busy":"2024-06-14T14:58:01.008202Z","iopub.execute_input":"2024-06-14T14:58:01.008727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Reinitialize the model and train.","metadata":{}},{"cell_type":"code","source":"dls = get_dls(trnds, vldds, bs=bs)\n\nmodel = get_model(vocab_size, n_embd, dropout)\niw = partial(init_weights, leaky=0.1)\nmodel = model.apply(iw)\n\nloss_func = partial(F.binary_cross_entropy_with_logits, pos_weight=pos_weight)\nmetrics = MetricsWithLogitsCB(BinaryAccuracy(), BinaryMatthewsCorrCoef(), BinaryAP(), BinaryAUROC())\ncbs = [DeviceCB(), ProgressCB(plot=False), metrics]\n\nlearn = BELKATrainLearner(model, dls, loss_func, lr=lr, cbs=cbs, opt_func=torch.optim.AdamW)\nprint(f\"Parameters total: {sum(p.nelement() for p in model.parameters())}\")\nlearn.fit(epochs)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 4.2 HSA Binding Classification","metadata":{}},{"cell_type":"markdown","source":"The next two models are trained using HSA binding data.","metadata":{}},{"cell_type":"markdown","source":"#### 4.2.1 Basic SMILES Tokenization","metadata":{}},{"cell_type":"code","source":"random_seed(seed_value, use_cuda=True)\nprint(f\"Random seed test: {torch.randint(0, 1000, (1,)).item()}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"xs, ys, valid_xs, valid_ys, vocab_size = get_train_valid_data(neg_size=1e6,\n                                                              bind_column='binds_HSA',\n                                                              tok_type='smiles')\ntrnds = Dataset(xs, ys)\nvldds = Dataset(valid_xs, valid_ys)\npos_weight = (ys.shape[0] - ys.sum()) / ys.sum()\nprint(f\"Positive weight: {pos_weight:.2f}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls = get_dls(trnds, vldds, bs=bs)\n\nmodel = get_model(vocab_size, n_embd, dropout)\niw = partial(init_weights, leaky=0.1)\nmodel = model.apply(iw)\n\nloss_func = partial(F.binary_cross_entropy_with_logits, pos_weight=pos_weight)\nmetrics = MetricsWithLogitsCB(BinaryAccuracy(), BinaryMatthewsCorrCoef(), BinaryAP(), BinaryAUROC())\ncbs = [DeviceCB(), ProgressCB(plot=False), metrics]\n\nlearn = BELKATrainLearner(model, dls, loss_func, lr=lr, cbs=cbs, opt_func=torch.optim.AdamW)\nprint(f\"Parameters total: {sum(p.nelement() for p in model.parameters())}\")\nlearn.fit(epochs)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### 4.2.2 AIS Tokenization","metadata":{}},{"cell_type":"code","source":"random_seed(seed_value, use_cuda=True)\nprint(f\"Random seed test: {torch.randint(0, 1000, (1,)).item()}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"xs, ys, valid_xs, valid_ys, vocab_size = get_train_valid_data(neg_size=1e6,\n                                                              bind_column='binds_HSA',\n                                                              tok_type='ais')\ntrnds = Dataset(xs, ys)\nvldds = Dataset(valid_xs, valid_ys)\npos_weight = (ys.shape[0] - ys.sum()) / ys.sum()\nprint(f\"Positive weight: {pos_weight:.2f}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls = get_dls(trnds, vldds, bs=bs)\n\nmodel = get_model(vocab_size, n_embd, dropout)\niw = partial(init_weights, leaky=0.1)\nmodel = model.apply(iw)\n\nloss_func = partial(F.binary_cross_entropy_with_logits, pos_weight=pos_weight)\nmetrics = MetricsWithLogitsCB(BinaryAccuracy(), BinaryMatthewsCorrCoef(), BinaryAP(), BinaryAUROC())\ncbs = [DeviceCB(), ProgressCB(plot=False), metrics]\n\nlearn = BELKATrainLearner(model, dls, loss_func, lr=lr, cbs=cbs, opt_func=torch.optim.AdamW)\nprint(f\"Parameters total: {sum(p.nelement() for p in model.parameters())}\")\nlearn.fit(epochs)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 4.3 sEH Binding Classification","metadata":{}},{"cell_type":"markdown","source":"Finally, two models are trained on sEH binding data.","metadata":{}},{"cell_type":"markdown","source":"#### 4.3.1 Basic SMILES Tokenization","metadata":{}},{"cell_type":"code","source":"random_seed(seed_value, use_cuda=True)\nprint(f\"Random seed test: {torch.randint(0, 1000, (1,)).item()}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"xs, ys, valid_xs, valid_ys, vocab_size = get_train_valid_data(neg_size=1e6,\n                                                              bind_column='binds_sEH',\n                                                              tok_type='smiles')\ntrnds = Dataset(xs, ys)\nvldds = Dataset(valid_xs, valid_ys)\npos_weight = (ys.shape[0] - ys.sum()) / ys.sum()\nprint(f\"Positive weight: {pos_weight:.2f}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls = get_dls(trnds, vldds, bs=bs)\n\nmodel = get_model(vocab_size, n_embd, dropout)\niw = partial(init_weights, leaky=0.1)\nmodel = model.apply(iw)\n\nloss_func = partial(F.binary_cross_entropy_with_logits, pos_weight=pos_weight)\nmetrics = MetricsWithLogitsCB(BinaryAccuracy(), BinaryMatthewsCorrCoef(), BinaryAP(), BinaryAUROC())\ncbs = [DeviceCB(), ProgressCB(plot=False), metrics]\n\nlearn = BELKATrainLearner(model, dls, loss_func, lr=lr, cbs=cbs, opt_func=torch.optim.AdamW)\nprint(f\"Parameters total: {sum(p.nelement() for p in model.parameters())}\")\nlearn.fit(epochs)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### 4.3.2 AIS Tokenization","metadata":{}},{"cell_type":"code","source":"random_seed(seed_value, use_cuda=True)\nprint(f\"Random seed test: {torch.randint(0, 1000, (1,)).item()}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"xs, ys, valid_xs, valid_ys, vocab_size = get_train_valid_data(neg_size=1e6,\n                                                              bind_column='binds_sEH',\n                                                              tok_type='ais')\ntrnds = Dataset(xs, ys)\nvldds = Dataset(valid_xs, valid_ys)\npos_weight = (ys.shape[0] - ys.sum()) / ys.sum()\nprint(f\"Positive weight: {pos_weight:.2f}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls = get_dls(trnds, vldds, bs=bs)\n\nmodel = get_model(vocab_size, n_embd, dropout)\niw = partial(init_weights, leaky=0.1)\nmodel = model.apply(iw)\n\nloss_func = partial(F.binary_cross_entropy_with_logits, pos_weight=pos_weight)\nmetrics = MetricsWithLogitsCB(BinaryAccuracy(), BinaryMatthewsCorrCoef(), BinaryAP(), BinaryAUROC())\ncbs = [DeviceCB(), ProgressCB(plot=False), metrics]\n\nlearn = BELKATrainLearner(model, dls, loss_func, lr=lr, cbs=cbs, opt_func=torch.optim.AdamW)\nprint(f\"Parameters total: {sum(p.nelement() for p in model.parameters())}\")\nlearn.fit(epochs)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5. Conclusion","metadata":{}},{"cell_type":"markdown","source":"The table below summarizes the results obtained from training the convolutional model on all three proteins in the [LeashBio BELKA](https://www.kaggle.com/competitions/leash-BELKA) dataset using basic SMILES and AIS tokenization.\n\n| Protein | AP Score SMILES | AP Score AIS | Percent Improvement |\n|:--------|----------------:|-------------:|--------------------:|\n| BRD4    |           0.194 |        0.290 |               49.5 %|\n| HSA     |           0.108 |        0.140 |               29.6 %|\n| sEH     |           0.527 |        0.599 |               13.7 %|\n\nThe direct comparison between SMILES and AIS tokenization reveals that the [average precision](https://lightning.ai/docs/torchmetrics/stable/classification/average_precision.html) score obtained by a sequential model for binary classification of protein binding of molecules in the [LeashBio BELKA](https://www.kaggle.com/competitions/leash-BELKA) dataset can be significantly improved by 14-50 % using AtomInSmiles istead of basic SMILES tokenization.","metadata":{}}]}