{"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":"markdown","source":"### References:\n\n- [starter notebook from Y. Nakama](https://www.kaggle.com/yasufuminakama/inchi-resnet-lstm-with-attention-starter)\n- [adapted notebook from Konrad](https://www.kaggle.com/yasufuminakama/inchi-resnet-lstm-with-attention-starter)\n- [PyTorch tutorial on image captioning](https://github.com/sgrvinod/a-PyTorch-Tutorial-to-Image-Captioning)\n- [two-layer RNN implementation](https://github.com/sgrvinod/a-PyTorch-Tutorial-to-Image-Captioning/pull/79)","metadata":{}},{"cell_type":"code","source":"import os\nfrom matplotlib import pyplot as plt\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n\nimport sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\n\nimport os\nimport gc\nimport re\nimport math\nimport time\nimport random\nimport shutil\nimport pickle\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport Levenshtein\nfrom sklearn import preprocessing\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\n\nfrom functools import partial\n\nimport cv2\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.nn.utils.rnn import pad_sequence, pack_padded_sequence\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\n\nfrom albumentations import (\n    Compose, OneOf, Normalize, Resize, RandomResizedCrop, RandomCrop, HorizontalFlip, VerticalFlip, \n    RandomBrightness, RandomContrast, RandomBrightnessContrast, Rotate, ShiftScaleRotate, Cutout, \n    IAAAdditiveGaussianNoise, Transpose, Blur\n    )\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\nimport timm\nimport pytorch_lightning as pl\n\n# ! pip install nltk==3.7\nfrom nltk.translate import bleu_score\nfrom nltk.metrics import distance, scores\n\nimport warnings \nwarnings.filterwarnings('ignore')\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"papermill":{"duration":5.105933,"end_time":"2021-04-07T08:27:02.380708","exception":false,"start_time":"2021-04-07T08:26:57.274775","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:04.853314Z","iopub.execute_input":"2022-10-26T16:36:04.853747Z","iopub.status.idle":"2022-10-26T16:36:07.166844Z","shell.execute_reply.started":"2022-10-26T16:36:04.853622Z","shell.execute_reply":"2022-10-26T16:36:07.165950Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"CFG class now includes a new parameter: `decoder_layers`. For illustration purposes, I am running a two-layer LSTM for 1 epoch on 100k images.","metadata":{}},{"cell_type":"code","source":"# print(timm.list_models(pretrained=True))","metadata":{"execution":{"iopub.status.busy":"2022-10-26T16:36:07.168859Z","iopub.execute_input":"2022-10-26T16:36:07.169204Z","iopub.status.idle":"2022-10-26T16:36:07.174683Z","shell.execute_reply.started":"2022-10-26T16:36:07.169152Z","shell.execute_reply":"2022-10-26T16:36:07.173854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#  n_channels_dict = {'efficientnet-b0': 1280, 'efficientnet-b1': 1280, 'efficientnet-b2': 1408,\n#   'efficientnet-b3': 1536, 'efficientnet-b4': 1792, 'efficientnet-b5': 2048,\n#   'efficientnet-b6': 2304, 'efficientnet-b7': 2560}\n\n# This is not, to put it mildly, the most elegant solution ever - but I ran into some trouble \n# with checking the size of feature spaces programmatically inside the CFG definition.\n\nclass CFG:\n    debug          = True\n    apex           = False\n    max_len        = 278\n    print_freq     = 250\n    num_workers    = 4\n    num_pixels     = 9\n#     model_name     = 'efficientnetv2_rw_m'\n#     enc_size       = 2152\n#     model_name     = 'efficientnetv2_rw_s'\n#     enc_size       = 1792\n    \n    model_name     = 'efficientnet_b0'\n    enc_size       = 1280\n#     model_name     = 'mobilenetv2_100'\n#     enc_size       = 1280\n#     model_name     = 'resnet50'\n#     enc_size       = 2048\n#     model_name     = 'tnt_s_patch16_224'\n#     enc_size       = 384\n#     model_name     = 'vit_base_patch16_224'\n#     enc_size       = 768\n    samp_size      = 10000\n    size           = 288\n#     size           = 224\n    scheduler      = 'ReduceLROnPlateau'\n#     scheduler      = 'CosineAnnealingLR' \n    epochs         = 60\n    T_max          = 4  \n    encoder_lr     = 1e-4\n    decoder_lr     = 4e-4\n    min_lr         = 1e-6\n    batch_size     = 32\n    weight_decay   = 1e-6\n    gradient_accumulation_steps = 1\n    max_grad_norm  = 10\n    attention_dim  = 256\n    embed_dim      = 512\n    decoder_dim    = 512\n    decoder_layers = 2     # number of LSTM layers\n    dropout        = 0.1\n    seed           = 42\n    n_fold         = 5\n    trn_fold       = 0 \n    train          = True\n    train_path     = '../input/bms-molecular-translation/'\n    prep_path      = '../input/preprocessed-stuff/'\n    prev_model     = '../input/molecular-translation/efficientnet_b2_fold0_best.pth'\n    pred_model     = '../input/efficientnet-multilayer-lstm-4-epochs/efficientnet_b2_fold0_best.pth'","metadata":{"papermill":{"duration":0.022485,"end_time":"2021-04-07T08:27:10.713561","exception":false,"start_time":"2021-04-07T08:27:10.691076","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:07.176361Z","iopub.execute_input":"2022-10-26T16:36:07.177142Z","iopub.status.idle":"2022-10-26T16:36:07.188402Z","shell.execute_reply.started":"2022-10-26T16:36:07.177000Z","shell.execute_reply":"2022-10-26T16:36:07.187735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions","metadata":{"papermill":{"duration":0.011683,"end_time":"2021-04-07T08:27:10.737552","exception":false,"start_time":"2021-04-07T08:27:10.725869","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class Tokenizer(object):\n    \n    def __init__(self):\n        self.stoi = {}\n        self.itos = {}\n\n    def __len__(self):\n        return len(self.stoi)\n    \n    def fit_on_texts(self, texts):\n        vocab = set()\n        for text in texts:\n            vocab.update(text.split(' '))\n        vocab = sorted(vocab)\n        vocab.append('<sos>')\n        vocab.append('<eos>')\n        vocab.append('<pad>')\n        for i, s in enumerate(vocab):\n            self.stoi[s] = i\n        self.itos = {item[1]: item[0] for item in self.stoi.items()}\n        \n    def text_to_sequence(self, text):\n        sequence = []\n        sequence.append(self.stoi['<sos>'])\n        for s in text.split(' '):\n            sequence.append(self.stoi[s])\n        sequence.append(self.stoi['<eos>'])\n        return sequence\n    \n    def texts_to_sequences(self, texts):\n        sequences = []\n        for text in texts:\n            sequence = self.text_to_sequence(text)\n            sequences.append(sequence)\n        return sequences\n\n    def sequence_to_text(self, sequence):\n        return ''.join(list(map(lambda i: self.itos[i], sequence)))\n    \n    def sequences_to_texts(self, sequences):\n        texts = []\n        for sequence in sequences:\n            text = self.sequence_to_text(sequence)\n            texts.append(text)\n        return texts\n    \n    def predict_caption(self, sequence):\n        caption = ''\n        for i in sequence:\n            if i == self.stoi['<eos>'] or i == self.stoi['<pad>']:\n                break\n            caption += self.itos[i]\n        return caption\n    \n    def predict_captions(self, sequences):\n        captions = []\n        for sequence in sequences:\n            caption = self.predict_caption(sequence)\n            captions.append(caption)\n        return captions\n\ntokenizer = torch.load(CFG.prep_path + 'tokenizer2.pth')\nprint(f\"tokenizer.stoi: {tokenizer.stoi}\")","metadata":{"papermill":{"duration":0.037989,"end_time":"2021-04-07T08:27:10.787242","exception":false,"start_time":"2021-04-07T08:27:10.749253","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:07.191538Z","iopub.execute_input":"2022-10-26T16:36:07.191788Z","iopub.status.idle":"2022-10-26T16:36:07.206440Z","shell.execute_reply.started":"2022-10-26T16:36:07.191764Z","shell.execute_reply":"2022-10-26T16:36:07.205107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_score(y_true, y_pred):\n    bleus = []\n    distances = []\n    distances_norm = []\n    accs = []\n    for true, pred in zip(y_true, y_pred):\n        pred = pred[:len(true)]\n#         score1 = Levenshtein.distance(true, pred)\n        dist = distance.edit_distance(true, pred)\n        dist_norm = distance.edit_distance(true, pred)/len(true)\n        bleu = bleu_score.sentence_bleu([list(true)], list(pred))\n        bleus.append(bleu)\n        distances.append(dist)\n        distances_norm.append(dist_norm)\n        if len(true) == len(pred):\n            acc = scores.accuracy(list(true), list(pred))\n            accs.append(acc)\n    avg_bleu = np.mean(bleus)\n    avg_dist = np.mean(distances)\n    avg_dist_norm = np.mean(distances_norm)\n    avg_acc = np.mean(accs)\n    return avg_bleu, avg_dist, avg_dist_norm, avg_acc\n\n\ndef init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\n\ndef seed_torch(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_torch(seed = CFG.seed)","metadata":{"papermill":{"duration":0.027729,"end_time":"2021-04-07T08:27:10.827442","exception":false,"start_time":"2021-04-07T08:27:10.799713","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:07.211799Z","iopub.execute_input":"2022-10-26T16:36:07.212039Z","iopub.status.idle":"2022-10-26T16:36:07.224154Z","shell.execute_reply.started":"2022-10-26T16:36:07.212015Z","shell.execute_reply":"2022-10-26T16:36:07.223432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\n\nclass TrainDataset(Dataset):\n    def __init__(self, df, tokenizer, transform=None):\n        super().__init__()\n        self.df         = df\n        self.tokenizer  = tokenizer\n        self.file_paths = df['file_path'].values\n        self.labels     = df['InChI_text'].values\n        self.transform  = transform\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        file_path = self.file_paths[idx]\n        image = cv2.imread(file_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n        if self.transform:\n            augmented = self.transform(image = image)\n            image     = augmented['image']\n        label = self.labels[idx]\n        label = self.tokenizer.text_to_sequence(label)\n        label_length = len(label)\n        label_length = torch.LongTensor([label_length])\n        return image, torch.LongTensor(label), label_length\n    \n\nclass TestDataset(Dataset):\n    def __init__(self, df, transform=None):\n        super().__init__()\n        self.df = df\n        self.file_paths = df['file_path'].values\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        file_path = self.file_paths[idx]\n        image = cv2.imread(file_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n        return image","metadata":{"papermill":{"duration":0.024936,"end_time":"2021-04-07T08:27:10.869493","exception":false,"start_time":"2021-04-07T08:27:10.844557","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:07.227295Z","iopub.execute_input":"2022-10-26T16:36:07.227757Z","iopub.status.idle":"2022-10-26T16:36:07.238717Z","shell.execute_reply.started":"2022-10-26T16:36:07.227723Z","shell.execute_reply":"2022-10-26T16:36:07.237723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def bms_collate(batch):\n    imgs, labels, label_lengths = [], [], []\n    for data_point in batch:\n        imgs.append(data_point[0])\n        labels.append(data_point[1])\n        label_lengths.append(data_point[2])\n    labels = pad_sequence(labels, batch_first = True, padding_value = tokenizer.stoi[\"<pad>\"])\n    return torch.stack(imgs), labels, torch.stack(label_lengths).reshape(-1, 1)","metadata":{"papermill":{"duration":0.021152,"end_time":"2021-04-07T08:27:10.902922","exception":false,"start_time":"2021-04-07T08:27:10.88177","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:07.240460Z","iopub.execute_input":"2022-10-26T16:36:07.240855Z","iopub.status.idle":"2022-10-26T16:36:07.249123Z","shell.execute_reply.started":"2022-10-26T16:36:07.240818Z","shell.execute_reply":"2022-10-26T16:36:07.248377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"####### CNN ENCODER\n\nclass Encoder(nn.Module):\n    def __init__(self, model_name = CFG.model_name, pretrained = False):\n        super().__init__()\n        self.cnn = timm.create_model(model_name, pretrained = pretrained)\n\n    def forward(self, x):\n        bs       = x.size(0)\n        features = self.cnn.forward_features(x)\n\n        features = features.permute(0, 2, 3, 1)\n\n        return features","metadata":{"papermill":{"duration":0.021209,"end_time":"2021-04-07T08:27:10.936685","exception":false,"start_time":"2021-04-07T08:27:10.915476","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:07.250448Z","iopub.execute_input":"2022-10-26T16:36:07.251014Z","iopub.status.idle":"2022-10-26T16:36:07.259140Z","shell.execute_reply.started":"2022-10-26T16:36:07.250980Z","shell.execute_reply":"2022-10-26T16:36:07.258419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The class `DecoderWithAttention` is updated to support a multi-layer LSTM.","metadata":{}},{"cell_type":"code","source":"####### RNN DECODER\n\n# attention module\nclass Attention(nn.Module):\n    '''\n    Attention network for calculate attention value\n    '''\n    def __init__(self, encoder_dim, decoder_dim, attention_dim):\n        '''\n        :param encoder_dim: input size of encoder network\n        :param decoder_dim: input size of decoder network\n        :param attention_dim: input size of attention network\n        '''\n        super(Attention, self).__init__()\n        self.encoder_att = nn.Linear(encoder_dim, attention_dim)  # linear layer to transform encoded image\n        self.decoder_att = nn.Linear(decoder_dim, attention_dim)  # linear layer to transform decoder's output\n        self.full_att    = nn.Linear(attention_dim, 1)            # linear layer to calculate values to be softmax-ed\n        self.relu        = nn.ReLU()\n        self.softmax     = nn.Softmax(dim = 1)  # softmax layer to calculate weights\n\n    def forward(self, encoder_out, decoder_hidden):\n        att1  = self.encoder_att(encoder_out)     # (batch_size, num_pixels, attention_dim)\n        att2  = self.decoder_att(decoder_hidden)  # (batch_size, attention_dim)\n        att   = self.full_att(self.relu(att1 + att2.unsqueeze(1))).squeeze(2)  # (batch_size, num_pixels)\n        alpha = self.softmax(att)                 # (batch_size, num_pixels)\n        attention_weighted_encoding = (encoder_out * alpha.unsqueeze(2)).sum(dim = 1)  # (batch_size, encoder_dim)\n        return attention_weighted_encoding, alpha\n    \n    \n# custom LSTM cell\ndef LSTMCell(input_size, hidden_size, **kwargs):\n    m = nn.LSTMCell(input_size, hidden_size, **kwargs)\n    for name, param in m.named_parameters():\n        if 'weight' in name or 'bias' in name:\n            param.data.uniform_(-0.1, 0.1)\n    return m\n\n\n# decoder\nclass DecoderWithAttention(nn.Module):\n    '''\n    Decoder network with attention network used for training\n    '''\n\n    def __init__(self, attention_dim, embed_dim, decoder_dim, vocab_size, device, encoder_dim, dropout, num_layers):\n        '''\n        :param attention_dim: input size of attention network\n        :param embed_dim: input size of embedding network\n        :param decoder_dim: input size of decoder network\n        :param vocab_size: total number of characters used in training\n        :param encoder_dim: input size of encoder network\n        :param num_layers: number of the LSTM layers\n        :param dropout: dropout rate\n        '''\n        super(DecoderWithAttention, self).__init__()\n        self.encoder_dim   = encoder_dim\n        self.attention_dim = attention_dim\n        self.embed_dim     = embed_dim\n        self.decoder_dim   = decoder_dim\n        self.vocab_size    = vocab_size\n        self.dropout       = dropout\n        self.num_layers    = num_layers\n        self.device        = device\n        self.attention     = Attention(encoder_dim, decoder_dim, attention_dim)  # attention network\n        self.embedding     = nn.Embedding(vocab_size, embed_dim)                 # embedding layer\n        self.dropout       = nn.Dropout(p = self.dropout)\n        self.decode_step   = nn.ModuleList([LSTMCell(embed_dim + encoder_dim if layer == 0 else embed_dim, embed_dim) for layer in range(self.num_layers)]) # decoding LSTMCell        \n        self.init_h        = nn.Linear(encoder_dim, decoder_dim)  # linear layer to find initial hidden state of LSTMCell\n        self.init_c        = nn.Linear(encoder_dim, decoder_dim)  # linear layer to find initial cell state of LSTMCell\n        self.f_beta        = nn.Linear(decoder_dim, encoder_dim)  # linear layer to create a sigmoid-activated gate\n        self.sigmoid       = nn.Sigmoid()\n        self.fc            = nn.Linear(decoder_dim, vocab_size)  # linear layer to find scores over vocabulary\n        self.init_weights()                                      # initialize some layers with the uniform distribution\n\n    def init_weights(self):\n        self.embedding.weight.data.uniform_(-0.1, 0.1)\n        self.fc.bias.data.fill_(0)\n        self.fc.weight.data.uniform_(-0.1, 0.1)\n\n    def load_pretrained_embeddings(self, embeddings):\n        self.embedding.weight = nn.Parameter(embeddings)\n\n    def fine_tune_embeddings(self, fine_tune = True):\n        for p in self.embedding.parameters():\n            p.requires_grad = fine_tune\n\n    def init_hidden_state(self, encoder_out):\n        mean_encoder_out = encoder_out.mean(dim = 1)\n        h = [self.init_h(mean_encoder_out) for i in range(self.num_layers)]  # (batch_size, decoder_dim)\n        c = [self.init_c(mean_encoder_out) for i in range(self.num_layers)]\n        return h, c\n\n    def forward(self, encoder_out, encoded_captions, decode_lengths):\n        '''\n        :param encoder_out: output of encoder network\n        :param encoded_captions: transformed sequence from character to integer\n        :param caption_lengths: length of transformed sequence\n        '''\n        batch_size       = encoder_out.size(0)\n        encoder_dim      = encoder_out.size(-1)\n        vocab_size       = self.vocab_size\n        encoder_out      = encoder_out.view(batch_size, -1, encoder_dim)  # (batch_size, num_pixels, encoder_dim)\n        num_pixels       = encoder_out.size(1)\n        \n#         caption_lengths, sort_ind = caption_lengths.squeeze(1).sort(dim = 0, descending = True)\n#         encoder_out      = encoder_out[sort_ind]\n#         encoded_captions = encoded_captions[sort_ind]\n        \n        \n        # embedding transformed sequence for vector\n        embeddings = self.embedding(encoded_captions)  # (batch_size, max_caption_length, embed_dim)\n        \n        # Initialize LSTM state, initialize cell_vector and hidden_vector\n        prev_h, prev_c = self.init_hidden_state(encoder_out)  # (batch_size, decoder_dim)\n        \n        # set decode length by caption length - 1 because of omitting start token\n#         decode_lengths = (caption_lengths - 1).tolist()\n        predictions    = torch.zeros(batch_size, max(decode_lengths), vocab_size, device = self.device)\n        alphas         = torch.zeros(batch_size, max(decode_lengths), num_pixels, device = self.device)\n        \n        # predict sequence\n        for t in range(max(decode_lengths)):\n            batch_size_t = sum([l > t for l in decode_lengths])\n            attention_weighted_encoding, alpha = self.attention(encoder_out[:batch_size_t],\n                                                                prev_h[-1][:batch_size_t])\n            gate = self.sigmoid(self.f_beta(prev_h[-1][:batch_size_t]))  # gating scalar, (batch_size_t, encoder_dim)\n            attention_weighted_encoding = gate * attention_weighted_encoding\n\n            input = torch.cat([embeddings[:batch_size_t, t, :], attention_weighted_encoding], dim=1)\n            \n            for i, rnn in enumerate(self.decode_step):\n                # recurrent cell\n                h, c = rnn(input, (prev_h[i][:batch_size_t], prev_c[i][:batch_size_t])) # cell_vector and hidden_vector\n\n                # hidden state becomes the input to the next layer\n                input = self.dropout(h)\n\n                # save state for next time step\n                prev_h[i] = h\n                prev_c[i] = c\n                \n            preds = self.fc(self.dropout(h))  # (batch_size_t, vocab_size)\n            predictions[:batch_size_t, t, :] = preds\n            alphas[:batch_size_t, t, :]      = alpha\n            \n        return predictions\n    \n    def predict(self, encoder_out, decode_lengths, tokenizer):\n        \n        # size variables\n        batch_size  = encoder_out.size(0)\n        encoder_dim = encoder_out.size(-1)\n        vocab_size  = self.vocab_size\n        encoder_out = encoder_out.view(batch_size, -1, encoder_dim)  # (batch_size, num_pixels, encoder_dim)\n        num_pixels  = encoder_out.size(1)\n        \n        # embed start tocken for LSTM input\n        start_tockens = torch.ones(batch_size, dtype = torch.long, device = self.device) * tokenizer.stoi['<sos>']\n        embeddings    = self.embedding(start_tockens)\n        \n        # initialize hidden state and cell state of LSTM cell\n        h, c        = self.init_hidden_state(encoder_out)  # (batch_size, decoder_dim)\n        predictions = torch.zeros(batch_size, decode_lengths, vocab_size, device = self.device)\n        \n        # predict sequence\n        end_condition = torch.zeros(batch_size, dtype=torch.long, device = self.device)\n        for t in range(decode_lengths):\n            awe, alpha = self.attention(encoder_out, h[-1])  # (s, encoder_dim), (s, num_pixels)\n            gate       = self.sigmoid(self.f_beta(h[-1]))    # gating scalar, (s, encoder_dim)\n            awe        = gate * awe\n            \n            input = torch.cat([embeddings, awe], dim=1)\n \n            for j, rnn in enumerate(self.decode_step):\n                at_h, at_c = rnn(input, (h[j], c[j]))  # (s, decoder_dim)\n                input = self.dropout(at_h)\n                h[j]  = at_h\n                c[j]  = at_c\n            \n            preds = self.fc(self.dropout(h[-1]))  # (batch_size_t, vocab_size)\n            predictions[:, t, :] = preds\n            end_condition |= (torch.argmax(preds, -1) == tokenizer.stoi[\"<eos>\"])\n            if end_condition.sum() == batch_size:\n                break\n            embeddings = self.embedding(torch.argmax(preds, -1))\n        \n        return predictions\n    \n    # beam search\n    def forward_step(self, prev_tokens, hidden, encoder_out, function):\n        \n        h, c = hidden\n        #h, c = h.squeeze(0), c.squeeze(0)\n        h, c = [hi.squeeze(0) for hi in h], [ci.squeeze(0) for ci in c]\n        \n        embeddings = self.embedding(prev_tokens)\n        if embeddings.dim() == 3:\n            embeddings = embeddings.squeeze(1)\n            \n        awe, alpha = self.attention(encoder_out, h[-1])  # (s, encoder_dim), (s, num_pixels)\n        gate       = self.sigmoid(self.f_beta(h[-1]))    # gating scalar, (s, encoder_dim)\n        awe        = gate * awe\n        \n        input = torch.cat([embeddings, awe], dim = 1)\n        for j, rnn in enumerate(self.decode_step):\n            at_h, at_c = rnn(input, (h[j], c[j]))  # (s, decoder_dim)\n            input = self.dropout(at_h)\n            h[j]  = at_h\n            c[j]  = at_c\n\n        preds = self.fc(self.dropout(h[-1]))  # (batch_size_t, vocab_size)\n\n        #hidden = (h.unsqueeze(0), c.unsqueeze(0))\n        hidden = [hi.unsqueeze(0) for hi in h], [ci.unsqueeze(0) for ci in c]\n        predicted_softmax = function(preds, dim = 1)\n        \n        return predicted_softmax, hidden, None","metadata":{"papermill":{"duration":0.067717,"end_time":"2021-04-07T08:27:11.017304","exception":false,"start_time":"2021-04-07T08:27:10.949587","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:07.261806Z","iopub.execute_input":"2022-10-26T16:36:07.262085Z","iopub.status.idle":"2022-10-26T16:36:07.295283Z","shell.execute_reply.started":"2022-10-26T16:36:07.262033Z","shell.execute_reply":"2022-10-26T16:36:07.294548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PatchEmbedding(nn.Module):\n    def __init__(self, img_size, patch_size, in_chans, embed_dim):\n        super().__init__()\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.n_patches = (img_size // patch_size) ** 2\n        self.patch_size = patch_size\n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n\n    def forward(self, x):\n        x = x.permute(0, 3, 1, 2)\n        x = self.proj(x)  # (B, E, P, P)\n        x = x.flatten(2)  # (B, E, N)\n        x = x.transpose(1, 2)  # (B, N, E)\n#         print(x.size())\n        return x\n    \nclass Transformer(pl.LightningModule):\n\n    def __init__(self, \n                 len_vocab,\n                 img_size=28, \n                 patch_size=7, \n                 in_chans=1, \n                 embed_dim=100, \n                 max_len=8, \n                 nhead=2, \n                 num_encoder_layers=3,\n                 num_decoder_layers=3,\n                 dim_feedforward=400,\n                 dropout=0.1\n                ):\n        super().__init__()\n        \n        self.patch_embed = PatchEmbedding(img_size, patch_size, in_chans, embed_dim)\n        self.pos_embed = nn.Parameter(torch.zeros(1, self.patch_embed.n_patches, embed_dim))\n#         self.pos_embed = nn.Parameter(torch.zeros(1, 197, embed_dim))\n\n\n        self.trg_emb = nn.Embedding(len_vocab, embed_dim)\n        self.trg_pos_emb = nn.Embedding(max_len, embed_dim)\n        self.max_len = max_len\n\n        self.transformer = torch.nn.Transformer(\n            embed_dim, nhead, num_encoder_layers, num_decoder_layers, dim_feedforward, dropout\n        )\n        \n        self.l = nn.LayerNorm(embed_dim)\n#         self.fc0 = nn.Linear(in_chans, embed_dim)\n        self.fc = nn.Linear(embed_dim, len_vocab)\n\n    def forward(self, images, captions, placeholder=None):\n\n        # embed images\n        embed_imgs = self.patch_embed(images)\n#         embed_imgs = self.fc0(images)\n#         embed_imgs = embed_imgs.view(embed_imgs.size(0), -1, embed_imgs.size(1))\n#         print(embed_imgs.size())\n#         print(self.pos_embed.size()  )\n        \n        embed_imgs = embed_imgs + self.pos_embed  \n        # embed captions\n        B, trg_seq_len = captions.shape\n        trg_positions = (torch.arange(0, trg_seq_len).expand(B, trg_seq_len).to(self.device))\n        embed_trg = self.trg_emb(captions) + self.trg_pos_emb(trg_positions)\n        trg_mask = self.transformer.generate_square_subsequent_mask(trg_seq_len).to(self.device)\n        tgt_padding_mask = captions == 192\n        # transformer\n        y = self.transformer(\n            embed_imgs.permute(1,0,2),  \n            embed_trg.permute(1,0,2),  \n            tgt_mask=trg_mask, \n            tgt_key_padding_mask = tgt_padding_mask\n        ).permute(1,0,2) \n        # head\n\n        pred = self.fc(self.l(y))\n        \n        return pred\n    \n\n    def predict(self, images, label_len=None):\n        \n        self.eval()\n        with torch.no_grad():\n            images = images.to(self.device)\n            B = images.shape[0]\n            eos = torch.tensor([190], dtype=torch.long, device=self.device).expand(B, 1)\n            trg_input = eos\n            \n            if label_len is None: \n                max_len = self.max_len\n            else:\n                max_len = label_len\n                \n            for _ in range(max_len):\n                preds = self(images, trg_input)\n#                 print(trg_input[0])\n#                 print('pred', preds[0, 0])\n                preds_max = torch.argmax(preds, axis=2)\n                # print(trg_input[0])\n                # print(preds_max[0])\n                # print('-----')\n                \n                trg_input = torch.cat([eos, preds_max], 1)\n        \n            return preds\n        \n    def compute_loss_and_acc(self, batch):\n        x, y = batch\n        y_hat = self(x, y[:,:-1])\n        trg_output = y[:,1:] \n        loss = F.cross_entropy(y_hat.permute(0,2,1), trg_output) \n        # I know this is not the best metric...\n        acc = (torch.argmax(y_hat, axis=2) == trg_output).sum().item() / (trg_output.shape[0]*trg_output.shape[1])\n        return loss, acc\n    \n    def training_step(self, batch, batch_idx):\n        loss, acc = self.compute_loss_and_acc(batch)\n        self.log('loss', loss)\n        self.log('acc', acc, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        loss, acc = self.compute_loss_and_acc(batch)\n        self.log('val_loss', loss, prog_bar=True)\n        self.log('val_acc', acc, prog_bar=True)\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=0.001)\n        return optimizer","metadata":{"execution":{"iopub.status.busy":"2022-10-26T16:36:07.296732Z","iopub.execute_input":"2022-10-26T16:36:07.297254Z","iopub.status.idle":"2022-10-26T16:36:07.317039Z","shell.execute_reply.started":"2022-10-26T16:36:07.297197Z","shell.execute_reply":"2022-10-26T16:36:07.316301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Helper functions\n\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val   = 0\n        self.avg   = 0\n        self.sum   = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val    = val\n        self.sum   += val * n\n        self.count += n\n        self.avg    = self.sum / self.count\n\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s   = now - since\n    es  = s / (percent)\n    rs  = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\n\ndef train_fn(train_loader, encoder, decoder, criterion, \n             encoder_optimizer, decoder_optimizer, epoch,\n             encoder_scheduler, decoder_scheduler, device):\n    \n    batch_time = AverageMeter()\n    data_time  = AverageMeter()\n    losses     = AverageMeter()\n    \n    # switch to train mode\n    encoder.train()\n    decoder.train()\n    \n    start = end = time.time()\n    global_step = 0\n    \n#     text_preds = []\n    \n    for step, (images, labels, label_lengths) in enumerate(train_loader):\n        \n        # measure data loading time\n        data_time.update(time.time() - end)\n        \n        images        = images.to(device)\n        labels        = labels.to(device)\n        label_lengths = label_lengths.to(device)\n        batch_size    = images.size(0)\n        \n        \n        features = encoder(images)\n#         print(features.size())\n#         asd\n        \n        label_lengths, sort_ind = label_lengths.squeeze(1).sort(dim=0, descending=True)\n        features = features[sort_ind]\n        caps_sorted = labels[sort_ind]\n        decode_lengths = (label_lengths - 1).tolist()\n        \n#         caps_sorted1 = caps_sorted[:, :-1]\n#         predictions = decoder(features, caps_sorted)\n        predictions = decoder(features, caps_sorted, decode_lengths)\n                \n        targets     = caps_sorted[:, 1:]\n\n        predictions = pack_padded_sequence(predictions, decode_lengths, batch_first=True).data\n        targets     = pack_padded_sequence(targets, decode_lengths, batch_first=True).data\n        loss        = criterion(predictions, targets)\n\n        \n        # record loss\n        losses.update(loss.item(), batch_size)\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n            \n        if CFG.apex:\n            with amp.scale_loss(loss, decoder_optimizer) as scaled_loss:\n                scaled_loss.backward()\n        else:\n            loss.backward()\n            \n        encoder_grad_norm = torch.nn.utils.clip_grad_norm_(encoder.parameters(), CFG.max_grad_norm)\n        decoder_grad_norm = torch.nn.utils.clip_grad_norm_(decoder.parameters(), CFG.max_grad_norm)\n        \n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            encoder_optimizer.step()\n            decoder_optimizer.step()\n            encoder_optimizer.zero_grad()\n            decoder_optimizer.zero_grad()\n            global_step += 1\n            \n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                  'Data {data_time.val:.3f} ({data_time.avg:.3f}) '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'Encoder Grad: {encoder_grad_norm:.4f}  '\n                  'Decoder Grad: {decoder_grad_norm:.4f}  '\n                  #'Encoder LR: {encoder_lr:.6f}  '\n                  #'Decoder LR: {decoder_lr:.6f}  '\n                  .format(\n                   epoch+1, step, len(train_loader), \n                   batch_time        = batch_time,\n                   data_time         = data_time, \n                   loss              = losses,\n                   remain            = timeSince(start, float(step+1)/len(train_loader)),\n                   encoder_grad_norm = encoder_grad_norm,\n                   decoder_grad_norm = decoder_grad_norm,\n                   #encoder_lr=encoder_scheduler.get_lr()[0],\n                   #decoder_lr=decoder_scheduler.get_lr()[0],\n                   ))\n                \n    return losses.avg\n\n\ndef valid_fn(valid_loader, encoder, decoder, tokenizer, criterion, device):\n    \n    batch_time = AverageMeter()\n    data_time  = AverageMeter()\n    losses     = AverageMeter()\n    \n    # switch to evaluation mode\n    encoder.eval()\n    decoder.eval()\n    \n    text_preds = []\n    start = end = time.time()\n    \n    for step, (images, labels, label_lengths) in enumerate(valid_loader):\n        \n        # measure data loading time\n        data_time.update(time.time() - end)\n        \n        images     = images.to(device)\n        labels        = labels.to(device)\n        label_lengths = label_lengths.to(device)\n        batch_size = images.size(0)\n        \n        with torch.no_grad():\n            features    = encoder(images)\n            \n            label_lengths, sort_ind = label_lengths.squeeze(1).sort(dim=0, descending=True)\n            features = features[sort_ind]\n            caps_sorted = labels[sort_ind]\n            decode_lengths = (label_lengths - 1).tolist()\n                        \n#             predictions = decoder.predict(features, CFG.max_len, tokenizer)\n\n#             caps_sorted1 = caps_sorted[:, :-1]\n#             predictions = decoder(features, caps_sorted)\n            predictions = decoder(features, caps_sorted, decode_lengths)\n            predictions_score = decoder.predict(features, labels.size(1)-1)\n\n        predicted_sequence = torch.argmax(predictions_score.detach().cpu(), -1).numpy()\n        _text_preds        = tokenizer.predict_captions(predicted_sequence)\n        text_preds.append(_text_preds)\n        \n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n            print('EVAL: [{0}/{1}] '\n                  'Data {data_time.val:.3f} ({data_time.avg:.3f}) '\n                  'Elapsed {remain:s} '\n                  .format(\n                   step, len(valid_loader), \n                   batch_time = batch_time,\n                   data_time  = data_time,\n                   remain     = timeSince(start, float(step+1)/len(valid_loader)),\n                   ))\n            \n        targets = caps_sorted[:, 1:]     \n\n        predictions = pack_padded_sequence(predictions, decode_lengths, batch_first=True).data\n        targets     = pack_padded_sequence(targets, decode_lengths, batch_first=True).data\n        \n        loss        = criterion(predictions, targets)\n        losses.update(loss.item(), batch_size)\n    \n    text_preds = np.concatenate(text_preds)\n    \n    return losses.avg, text_preds","metadata":{"papermill":{"duration":0.039205,"end_time":"2021-04-07T08:27:11.070219","exception":false,"start_time":"2021-04-07T08:27:11.031014","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:07.318748Z","iopub.execute_input":"2022-10-26T16:36:07.319199Z","iopub.status.idle":"2022-10-26T16:36:07.342821Z","shell.execute_reply.started":"2022-10-26T16:36:07.319128Z","shell.execute_reply":"2022-10-26T16:36:07.341725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Train loop\n# ====================================================\nload_prev = False\nif load_prev:\n    states = torch.load(CFG.prev_model,  map_location=torch.device('cpu'))\n\nencoder = Encoder(CFG.model_name, pretrained = not load_prev)\nif load_prev:\n    encoder.load_state_dict(states['encoder'])\nencoder.to(device)\n# decoder = DecoderWithAttention(attention_dim = CFG.attention_dim, \n#                                embed_dim     = CFG.embed_dim, \n#                                encoder_dim   = CFG.enc_size,\n#                                decoder_dim   = CFG.decoder_dim,\n#                                num_layers    = CFG.decoder_layers,\n#                                vocab_size    = len(tokenizer), \n#                                dropout       = CFG.dropout, \n#                                device        = device)\n\ndecoder = Transformer(\n                 len_vocab=len(tokenizer),\n                 img_size=CFG.num_pixels, \n                 patch_size=1, \n                 in_chans=CFG.enc_size, \n                 embed_dim=CFG.embed_dim, \n                 max_len=CFG.max_len, \n                 nhead=8, \n                 num_encoder_layers=3,\n                 num_decoder_layers=3,\n                 dim_feedforward=512,\n                 dropout=CFG.dropout\n                )\n# decoder = Transformer(\n#                  len_vocab=len(tokenizer),\n#                  img_size=CFG.num_pixels, \n#                  patch_size=1, \n#                  in_chans=CFG.enc_size, \n#                  embed_dim=512, \n#                  max_len=CFG.max_len, \n#                  nhead=8, \n#                  num_encoder_layers=6,\n#                  num_decoder_layers=6,\n#                  dim_feedforward=2048,\n#                  dropout=0.1\n#                 )\n\nif load_prev:\n    decoder.load_state_dict(states['decoder'])\ndecoder.to(device)\n\n\ndef train_loop(folds, fold, encoder, decoder):\n\n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n\n    # ====================================================\n    # loader\n    # ====================================================\n    trn_idx = folds[folds['fold'] != fold].index\n    val_idx0 = folds[folds['fold'] == fold].index\n    val_idx = val_idx0[:len(val_idx0)//2]\n    test_idx = val_idx0[len(val_idx0)//2:]\n    \n\n    train_folds  = folds.loc[trn_idx].reset_index(drop = True)\n    valid_folds  = folds.loc[val_idx].reset_index(drop = True)\n    valid_labels = valid_folds['InChI'].values\n    test_folds  = folds.loc[test_idx].reset_index(drop = True)\n    test_labels = test_folds['InChI'].values\n\n    train_dataset = TrainDataset(train_folds, tokenizer, transform = get_transforms(data = 'train'))\n#     valid_dataset = TestDataset(valid_folds, transform = get_transforms(data = 'valid'))\n    valid_dataset = TrainDataset(valid_folds, tokenizer, transform = get_transforms(data = 'valid'))\n    test_dataset = TrainDataset(test_folds, tokenizer, transform = get_transforms(data = 'valid'))\n\n\n    \n    train_loader = DataLoader(train_dataset, \n                              batch_size  = CFG.batch_size, \n                              shuffle     = True, \n                              num_workers = CFG.num_workers, \n                              pin_memory  = True,\n                              drop_last   = True, \n                              collate_fn  = bms_collate)\n#     valid_loader = DataLoader(valid_dataset, \n#                               batch_size  = CFG.batch_size, \n#                               shuffle     = False, \n#                               num_workers = CFG.num_workers,\n#                               pin_memory  = True, \n#                               drop_last   = False)\n    \n    valid_loader = DataLoader(valid_dataset, \n                              batch_size  = CFG.batch_size, \n                              shuffle     = False, \n                              num_workers = CFG.num_workers, \n                              pin_memory  = True,\n                              drop_last   = False, \n                              collate_fn  = bms_collate)\n    test_loader = DataLoader(test_dataset, \n                              batch_size  = CFG.batch_size, \n                              shuffle     = False, \n                              num_workers = CFG.num_workers, \n                              pin_memory  = True,\n                              drop_last   = False, \n                              collate_fn  = bms_collate)\n    \n    # ====================================================\n    # scheduler \n    # ====================================================\n    def get_scheduler(optimizer):\n        if CFG.scheduler=='ReduceLROnPlateau':\n            scheduler = ReduceLROnPlateau(optimizer, \n                                          mode     = 'min', \n#                                           factor   = CFG.factor, \n#                                           patience = CFG.patience, \n                                          verbose  = True, \n#                                           eps      = CFG.eps\n                                         )\n        elif CFG.scheduler=='CosineAnnealingLR':\n            scheduler = CosineAnnealingLR(optimizer, \n                                          T_max      = CFG.T_max, \n                                          eta_min    = CFG.min_lr, \n                                          last_epoch = -1)\n        elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n            scheduler = CosineAnnealingWarmRestarts(optimizer, \n                                                    T_0        = CFG.T_0, \n                                                    T_mult     = 1, \n                                                    eta_min    = CFG.min_lr, \n                                                    last_epoch = -1)\n        return scheduler\n\n    # ====================================================\n    # model & optimizer\n    # ====================================================\n\n    encoder_optimizer = Adam(encoder.parameters(), \n                             lr           = CFG.encoder_lr, \n                             weight_decay = CFG.weight_decay, \n                             amsgrad      = False)\n    if load_prev:\n        encoder_optimizer.load_state_dict(states['encoder_optimizer'])\n    encoder_scheduler = get_scheduler(encoder_optimizer)\n    if load_prev:\n        encoder_scheduler.load_state_dict(states['encoder_scheduler'])\n    \n    decoder_optimizer = Adam(decoder.parameters(), \n                             lr           = CFG.decoder_lr, \n                             weight_decay = CFG.weight_decay, \n                             amsgrad      = False)\n    if load_prev:\n        decoder_optimizer.load_state_dict(states['decoder_optimizer'])\n\n    decoder_scheduler = get_scheduler(decoder_optimizer)\n    if load_prev:\n        decoder_scheduler.load_state_dict(states['decoder_scheduler'])\n\n    # ====================================================\n    # loop\n    # ====================================================\n    criterion = nn.CrossEntropyLoss(ignore_index = tokenizer.stoi[\"<pad>\"])\n\n    best_score = np.inf\n    best_loss  = np.inf\n    if load_prev:\n        record_score = states['record_score']\n        record_score_norm = states['record_score_norm']\n        record_bleu = states['record_bleu']\n        record_acc = states['record_acc']\n        record_loss = states['record_loss']\n        record_loss_val = states['record_loss_val']\n    else:\n        record_score = []\n        record_score_norm = []\n        record_bleu = []\n        record_acc = []\n        record_loss = []\n        record_loss_val = []\n        record_loss_test = []\n        \n    times = []\n    \n    for epoch in range(CFG.epochs):\n        \n        start_time = time.time()\n        \n        # train\n        avg_loss = train_fn(train_loader, encoder, decoder, criterion, \n                            encoder_optimizer, decoder_optimizer, epoch, \n                            encoder_scheduler, decoder_scheduler, device)\n\n        # eval\n        avg_loss_val, text_preds = valid_fn(valid_loader, encoder, decoder, tokenizer, criterion, device)\n        text_preds = [f\"InChI=1S/{text}\" for text in text_preds]\n        avg_loss_test, text_preds_test = valid_fn(test_loader, encoder, decoder, tokenizer, criterion, device)\n\n#         text_preds_test = [f\"InChI=1S/{text}\" for text in text_preds_test]\n#         LOGGER.info(f\"val labels: {valid_labels[:5]}\")\n#         LOGGER.info(f\"val preds: {text_preds[:5]}\")\n        \n        # scoring\n        bleu, score, score_norm, acc = get_score(valid_labels, text_preds)        \n#         bleu2, score2, score_norm2, acc2 = get_score(test_labels, text_preds_test)        \n        \n        \n        if isinstance(encoder_scheduler, ReduceLROnPlateau):\n            encoder_scheduler.step(score_norm) # score\n        elif isinstance(encoder_scheduler, CosineAnnealingLR):\n            encoder_scheduler.step()\n        elif isinstance(encoder_scheduler, CosineAnnealingWarmRestarts):\n            encoder_scheduler.step()\n            \n        if isinstance(decoder_scheduler, ReduceLROnPlateau):\n            decoder_scheduler.step(score_norm) # score\n        elif isinstance(decoder_scheduler, CosineAnnealingLR):\n            decoder_scheduler.step()\n        elif isinstance(decoder_scheduler, CosineAnnealingWarmRestarts):\n            decoder_scheduler.step()\n\n        elapsed = time.time() - start_time\n        times.append(elapsed)\n\n#         LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  time: {elapsed:.0f}s')\n#         LOGGER.info(f'Epoch {epoch+1} - avg_val_loss: {avg_loss_val:.4f}')\n#         LOGGER.info(f'Epoch {epoch+1} - Score norm: {score_norm:.4f}')\n#         LOGGER.info(f'Epoch {epoch+1} - Score: {score:.4f}')\n#         LOGGER.info(f'Epoch {epoch+1} - Bleu: {bleu:.4f}')\n#         LOGGER.info(f'Epoch {epoch+1} - Acc: {acc:.4f}')\n        \n        record_score.append(round(score, 4))\n        record_score_norm.append(round(score_norm, 4))\n        record_bleu.append(round(bleu, 4))\n        record_acc.append(round(acc, 4))\n        record_loss.append(round(avg_loss, 4))\n        record_loss_val.append(round(avg_loss_val, 4))\n        record_loss_test.append(round(avg_loss_test, 4))\n\n        \n#         if score < best_score:\n        if True:\n            best_score = score\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Score: {best_score:.4f} Model')\n            torch.save({'encoder': encoder.state_dict(), \n                        'encoder_optimizer': encoder_optimizer.state_dict(), \n                        'encoder_scheduler': encoder_scheduler.state_dict(), \n                        'decoder': decoder.state_dict(), \n                        'decoder_optimizer': decoder_optimizer.state_dict(), \n                        'decoder_scheduler': decoder_scheduler.state_dict(), \n                        'text_preds': text_preds,\n                        'record_score': record_score,\n                        'record_score_norm': record_score_norm,\n                        'record_bleu': record_bleu,\n                        'record_acc': record_acc,\n                        'record_loss': record_loss,\n                        'record_loss_val': record_loss_val,\n                       },\n                        OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best.pth')\n            \n#         if epoch % 10 == 0:\n        LOGGER.info(f'train_loss = {record_loss}')\n        LOGGER.info(f'val_loss = {record_loss_val}')\n        LOGGER.info(f'test_loss = {record_loss_test}')\n        LOGGER.info(f'score = {record_score}')\n        LOGGER.info(f'score_norm = {record_score_norm}')\n        LOGGER.info(f'bleu = {record_bleu}')\n        LOGGER.info(f'acc = {record_acc}')\n        LOGGER.info(f'avg_time = {round(sum(times)/len(times))}')\n#             print(\"train_loss = \", record_loss)\n#             print(\"val_loss = \", record_loss_val)\n#             print(\"score = \", record_score)\n#             print(\"score_norm = \", record_score_norm)\n#             print(\"bleu = \", record_bleu)\n#             print(\"acc = \", record_acc)\n#             print('avg_time = ', round(sum(times)/len(times)))","metadata":{"papermill":{"duration":0.038367,"end_time":"2021-04-07T08:27:11.123364","exception":false,"start_time":"2021-04-07T08:27:11.084997","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:07.344201Z","iopub.execute_input":"2022-10-26T16:36:07.344763Z","iopub.status.idle":"2022-10-26T16:36:09.516169Z","shell.execute_reply.started":"2022-10-26T16:36:07.344728Z","shell.execute_reply":"2022-10-26T16:36:09.515279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_file_path(image_id):\n\n    return CFG.train_path + \"train/{}/{}/{}/{}.png\".format(\n        image_id[0], image_id[1], image_id[2], image_id \n    )\n\ndef get_test_file_path(image_id):\n\n    return CFG.train_path + \"test/{}/{}/{}/{}.png\".format(\n        image_id[0], image_id[1], image_id[2], image_id \n    )","metadata":{"papermill":{"duration":0.020522,"end_time":"2021-04-07T08:27:11.156788","exception":false,"start_time":"2021-04-07T08:27:11.136266","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:09.518994Z","iopub.execute_input":"2022-10-26T16:36:09.519365Z","iopub.status.idle":"2022-10-26T16:36:09.524350Z","shell.execute_reply.started":"2022-10-26T16:36:09.519335Z","shell.execute_reply":"2022-10-26T16:36:09.523470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# transformations\n\ndef get_transforms(*, data):\n    \n    if data == 'train':\n        return Compose([\n            Resize(CFG.size, CFG.size),\n            HorizontalFlip(p=0.5),                  \n            Transpose(p=0.5),\n            HorizontalFlip(p=0.5),\n            VerticalFlip(p=0.5),\n            ShiftScaleRotate(p=0.5),   \n            Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n    \n    elif data == 'valid':\n        return Compose([\n            Resize(CFG.size, CFG.size),\n            Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n","metadata":{"papermill":{"duration":0.022257,"end_time":"2021-04-07T08:27:11.192091","exception":false,"start_time":"2021-04-07T08:27:11.169834","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:09.525917Z","iopub.execute_input":"2022-10-26T16:36:09.526449Z","iopub.status.idle":"2022-10-26T16:36:09.536459Z","shell.execute_reply.started":"2022-10-26T16:36:09.526408Z","shell.execute_reply":"2022-10-26T16:36:09.535597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data","metadata":{"papermill":{"duration":0.013404,"end_time":"2021-04-07T08:27:11.218463","exception":false,"start_time":"2021-04-07T08:27:11.205059","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train = pd.read_pickle(CFG.prep_path + 'train2.pkl')\n\ntrain['file_path'] = train['image_id'].apply(get_train_file_path)\n\nprint(f'train.shape: {train.shape}')\n\ntest = pd.read_csv('../input/bms-molecular-translation/sample_submission.csv')\n\ntest['file_path'] = test['image_id'].apply(get_test_file_path)\n\nprint(f'test.shape: {test.shape}')\n\n\nif CFG.debug:\n    # CFG.epochs = 1\n    train = train.sample(n = CFG.samp_size, random_state = CFG.seed).reset_index(drop = True)","metadata":{"papermill":{"duration":12.663809,"end_time":"2021-04-07T08:27:23.895404","exception":false,"start_time":"2021-04-07T08:27:11.231595","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:09.537597Z","iopub.execute_input":"2022-10-26T16:36:09.537979Z","iopub.status.idle":"2022-10-26T16:36:16.960494Z","shell.execute_reply.started":"2022-10-26T16:36:09.537922Z","shell.execute_reply":"2022-10-26T16:36:16.959686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = TrainDataset(train, tokenizer, transform = get_transforms(data='train'))\n\nfolds = train.copy()\nFold = StratifiedKFold(n_splits = CFG.n_fold, shuffle = True, random_state = CFG.seed)\nfor n, (train_index, val_index) in enumerate(Fold.split(folds, folds['InChI_length'])):\n    folds.loc[val_index, 'fold'] = int(n)\nfolds['fold'] = folds['fold'].astype(int)","metadata":{"papermill":{"duration":0.02079,"end_time":"2021-04-07T08:27:23.931533","exception":false,"start_time":"2021-04-07T08:27:23.910743","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:16.961815Z","iopub.execute_input":"2022-10-26T16:36:16.962140Z","iopub.status.idle":"2022-10-26T16:36:16.982461Z","shell.execute_reply.started":"2022-10-26T16:36:16.962107Z","shell.execute_reply":"2022-10-26T16:36:16.981811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"p = sum(p.numel() for p in decoder.parameters())\npt = sum(p.numel() for p in decoder.parameters() if p.requires_grad)\nprint(p, pt)\n# transformer 15594689\n# lstm 9449666","metadata":{"execution":{"iopub.status.busy":"2022-10-26T16:36:16.983616Z","iopub.execute_input":"2022-10-26T16:36:16.983939Z","iopub.status.idle":"2022-10-26T16:36:16.991693Z","shell.execute_reply.started":"2022-10-26T16:36:16.983901Z","shell.execute_reply":"2022-10-26T16:36:16.990773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{"papermill":{"duration":0.013921,"end_time":"2021-04-07T08:27:24.004031","exception":false,"start_time":"2021-04-07T08:27:23.99011","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_loop(folds, CFG.trn_fold, encoder, decoder)","metadata":{"papermill":{"duration":39.171513,"end_time":"2021-04-07T08:28:03.18932","exception":false,"start_time":"2021-04-07T08:27:24.017807","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-10-26T16:36:16.993059Z","iopub.execute_input":"2022-10-26T16:36:16.993603Z","iopub.status.idle":"2022-10-26T17:07:18.400997Z","shell.execute_reply.started":"2022-10-26T16:36:16.993568Z","shell.execute_reply":"2022-10-26T17:07:18.398586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# b0 + trans patch 1\ntrain_loss1 = [2.1337, 1.5656, 1.4038, 1.2849, 1.1938, 1.1283, 1.0691, 1.0152, 0.964, 0.9224, 0.8839, 0.8506, 0.8171, 0.7849, 0.759, 0.7303, 0.7069, 0.6901, 0.6633, 0.6428, 0.624, 0.6033, 0.5897, 0.5706, 0.5557, 0.5393, 0.5347, 0.5197, 0.5058, 0.4875, 0.4764, 0.4754, 0.4724, 0.4596, 0.4464, 0.4308, 0.418, 0.4129, 0.4052, 0.3976, 0.3855, 0.3795, 0.3716, 0.3612, 0.3521, 0.353, 0.344, 0.2752, 0.2398, 0.2251, 0.2149, 0.2079, 0.2013, 0.1963, 0.1895, 0.1853, 0.1817, 0.1786, 0.1752, 0.1684, 0.1668]\nval_loss1 = [1.6277, 1.4302, 1.2829, 1.2219, 1.1287, 1.0591, 1.0389, 0.9552, 0.9193, 0.8973, 0.8482, 0.8251, 0.7998, 0.7783, 0.7596, 0.7465, 0.7398, 0.7258, 0.6911, 0.6905, 0.6926, 0.6611, 0.6696, 0.6493, 0.6447, 0.6383, 0.6369, 0.647, 0.6321, 0.6244, 0.6227, 0.658, 0.6356, 0.6456, 0.6321, 0.6205, 0.6454, 0.6318, 0.6341, 0.6353, 0.6342, 0.6472, 0.6374, 0.6596, 0.6401, 0.6588, 0.6493, 0.6203, 0.6232, 0.6248, 0.6334, 0.6364, 0.6407, 0.6461, 0.6538, 0.6577, 0.6602, 0.6699, 0.6719, 0.6703, 0.6711]\nscore_norm1 = [0.6261, 0.6153, 0.6124, 0.616, 0.601, 0.6071, 0.6072, 0.6021, 0.6037, 0.6007, 0.5972, 0.5997, 0.6002, 0.5986, 0.5968, 0.5985, 0.5941, 0.5968, 0.5953, 0.5961, 0.5957, 0.5966, 0.5963, 0.5966, 0.5944, 0.5934, 0.5956, 0.5937, 0.5947, 0.5939, 0.5941, 0.5953, 0.594, 0.5944, 0.594, 0.593, 0.5943, 0.5946, 0.5931, 0.594, 0.5937, 0.5941, 0.5928, 0.5923, 0.591, 0.5916, 0.5917, 0.5913, 0.5916, 0.5914, 0.5917, 0.5916, 0.5915, 0.591, 0.5913, 0.5907, 0.591, 0.5912, 0.591, 0.5909, 0.5909]\nbleu1 = [0.2786, 0.2953, 0.294, 0.2929, 0.318, 0.3145, 0.3158, 0.3262, 0.3264, 0.3237, 0.3324, 0.3348, 0.331, 0.3344, 0.3328, 0.3324, 0.338, 0.335, 0.3393, 0.3396, 0.3382, 0.3382, 0.3404, 0.3399, 0.3426, 0.3405, 0.3396, 0.3441, 0.3415, 0.3419, 0.3418, 0.3426, 0.3426, 0.3408, 0.3434, 0.3438, 0.3448, 0.3419, 0.3453, 0.3453, 0.3432, 0.3437, 0.3452, 0.3457, 0.3464, 0.3471, 0.3465, 0.3482, 0.3489, 0.3485, 0.3486, 0.3489, 0.349, 0.349, 0.3483, 0.3488, 0.3488, 0.3491, 0.3487, 0.3489, 0.349]\nacc1 = [0.2255, 0.2195, 0.2209, 0.2202, 0.2241, 0.2193, 0.2215, 0.2228, 0.2202, 0.225, 0.2231, 0.2277, 0.2267, 0.224, 0.2298, 0.2278, 0.2291, 0.2296, 0.2268, 0.2278, 0.225, 0.2247, 0.2261, 0.2307, 0.2301, 0.2309, 0.2275, 0.2289, 0.229, 0.2287, 0.2284, 0.2275, 0.2294, 0.2285, 0.2311, 0.2336, 0.2314, 0.2268, 0.2352, 0.2312, 0.2305, 0.2294, 0.2337, 0.2328, 0.2326, 0.2298, 0.2323, 0.2367, 0.2351, 0.2366, 0.2377, 0.2392, 0.237, 0.2376, 0.238, 0.2376, 0.2364, 0.2347, 0.238, 0.2399, 0.2398]\navg_time = 142\n\n#b0 + lstm\ntrain_loss2 = [2.2334, 1.5415, 1.4207, 1.3193, 1.2323, 1.15, 1.0863, 1.0275, 0.9763, 0.9322, 0.8884, 0.8526, 0.8294, 0.7887, 0.7604, 0.7293, 0.7048, 0.6779, 0.6533, 0.6296, 0.6056, 0.5944, 0.563, 0.5412, 0.5203, 0.4992, 0.4818, 0.4673, 0.4471, 0.4305, 0.4168, 0.3988, 0.3828, 0.3244, 0.3049, 0.2979, 0.2926, 0.2881, 0.2833, 0.2781, 0.2757, 0.2712, 0.2681, 0.2658, 0.259, 0.2587, 0.2584, 0.2569, 0.2577, 0.2565, 0.2559, 0.2549, 0.2544, 0.2544, 0.2538, 0.2537, 0.2533, 0.2529, 0.2532, 0.2532, 0.2532]\nval_loss2 = [1.613, 1.6876, 1.4251, 1.2886, 1.1684, 1.1465, 1.0729, 1.0348, 0.9787, 0.9349, 0.9197, 0.8894, 0.8611, 0.8531, 0.8264, 0.8161, 0.8177, 0.7947, 0.7883, 0.788, 0.7868, 0.7834, 0.8024, 0.7857, 0.7891, 0.8046, 0.8001, 0.8024, 0.8213, 0.8208, 0.8477, 0.8471, 0.8531, 0.8348, 0.8407, 0.8465, 0.8523, 0.8569, 0.8632, 0.8693, 0.8738, 0.878, 0.8829, 0.8869, 0.8873, 0.8875, 0.888, 0.8882, 0.8888, 0.8902, 0.8902, 0.8928, 0.892, 0.8915, 0.8931, 0.8926, 0.8933, 0.8935, 0.8924, 0.8926, 0.8926]\nscore_norm2 = [0.6246, 0.6247, 0.6326, 0.6184, 0.6105, 0.6224, 0.6046, 0.6058, 0.6083, 0.6002, 0.5979, 0.6007, 0.5998, 0.5996, 0.5981, 0.5959, 0.6027, 0.5975, 0.5973, 0.5976, 0.5959, 0.5967, 0.5952, 0.5977, 0.5973, 0.5959, 0.5987, 0.5965, 0.5979, 0.5968, 0.596, 0.5975, 0.5968, 0.5953, 0.5954, 0.595, 0.5952, 0.5958, 0.5959, 0.5953, 0.5957, 0.5961, 0.5959, 0.5956, 0.5954, 0.5957, 0.5955, 0.5956, 0.5955, 0.5954, 0.5955, 0.5956, 0.5959, 0.5955, 0.5958, 0.5956, 0.5959, 0.5955, 0.5958, 0.5959, 0.5957]\nbleu2 = [0.3041, 0.3091, 0.3155, 0.3181, 0.3239, 0.3222, 0.3244, 0.3313, 0.3284, 0.3331, 0.3336, 0.3284, 0.3349, 0.3359, 0.3374, 0.3386, 0.3383, 0.3392, 0.336, 0.3386, 0.3407, 0.3392, 0.3399, 0.3377, 0.3389, 0.34, 0.3382, 0.34, 0.3395, 0.341, 0.3405, 0.339, 0.3397, 0.342, 0.3416, 0.3418, 0.3422, 0.342, 0.3416, 0.3422, 0.3421, 0.3414, 0.3417, 0.3417, 0.3419, 0.3418, 0.342, 0.3421, 0.3421, 0.342, 0.342, 0.3417, 0.3418, 0.342, 0.3422, 0.3418, 0.3417, 0.3421, 0.342, 0.3417, 0.342]\nacc2 = [0.2149, 0.2132, 0.2087, 0.2136, 0.2147, 0.2116, 0.221, 0.2211, 0.2198, 0.2271, 0.2295, 0.2266, 0.2268, 0.2244, 0.2242, 0.2289, 0.2223, 0.2262, 0.2239, 0.225, 0.2293, 0.2286, 0.2271, 0.2209, 0.2258, 0.2246, 0.2252, 0.2246, 0.2252, 0.2281, 0.2256, 0.2228, 0.2248, 0.2285, 0.2283, 0.2291, 0.2296, 0.2281, 0.2256, 0.2286, 0.2284, 0.226, 0.2282, 0.2262, 0.2267, 0.2269, 0.2268, 0.2266, 0.2266, 0.2278, 0.2262, 0.2265, 0.2264, 0.2282, 0.2272, 0.2278, 0.2263, 0.2271, 0.2266, 0.2274, 0.2273]\navg_time = 240\n\n#tnt\ntrain_loss3 = [2.1166, 1.5375, 1.3953, 1.3031, 1.2359, 1.182, 1.1409, 1.1065, 1.0756, 1.0479, 1.025, 1.0035, 0.9823, 0.9648, 0.9464, 0.929, 0.9149, 0.9002, 0.8848, 0.8706, 0.859, 0.845, 0.8329, 0.8222, 0.8114, 0.8005, 0.789, 0.78, 0.7702, 0.7613, 0.752, 0.743, 0.7338, 0.7273, 0.7193, 0.7105, 0.7048, 0.6973, 0.6255, 0.5993, 0.5884, 0.58, 0.5721, 0.568, 0.5629, 0.5586, 0.5552, 0.5513, 0.5476, 0.538, 0.5363, 0.5349, 0.535, 0.5362, 0.5341, 0.5337, 0.5335, 0.5331, 0.533, 0.5319, 0.5315]\nval_loss3 = [1.6007, 1.4041, 1.2945, 1.2194, 1.1722, 1.1244, 1.1087, 1.0795, 1.0469, 1.032, 1.0156, 1.007, 1.0024, 0.9913, 0.9734, 0.9731, 0.9623, 0.9534, 0.9446, 0.9506, 0.9489, 0.9391, 0.9455, 0.9433, 0.9424, 0.9448, 0.9374, 0.9522, 0.9531, 0.9533, 0.954, 0.9646, 0.9684, 0.9657, 0.9639, 0.9706, 0.9799, 0.9892, 0.9783, 0.9849, 0.9928, 0.9988, 1.0036, 1.0091, 1.0128, 1.0186, 1.0208, 1.0266, 1.0303, 1.0309, 1.0316, 1.0327, 1.0333, 1.0336, 1.0346, 1.0352, 1.0353, 1.0354, 1.0365, 1.0368, 1.0369]\nscore3 = [79.381, 78.4485, 77.9975, 76.587, 76.5785, 76.481, 75.9715, 76.109, 77.0015, 75.9045, 76.2225, 76.4745, 76.74, 76.323, 75.8705, 76.557, 76.396, 76.4885, 76.0325, 75.812, 76.8685, 76.08, 76.323, 76.123, 76.266, 75.936, 75.8855, 76.28, 76.7255, 76.178, 76.0955, 76.0515, 76.275, 76.19, 76.0715, 76.291, 76.4705, 76.2985, 75.94, 75.937, 75.8805, 75.976, 76.0375, 75.9525, 76.1275, 76.191, 76.0885, 76.047, 76.132, 76.056, 76.0945, 76.0785, 76.0715, 76.097, 76.1205, 76.087, 76.0565, 76.082, 76.056, 76.1015, 76.0805]\nscore_norm3 = [0.6273, 0.6202, 0.6169, 0.6056, 0.6057, 0.6047, 0.6008, 0.6016, 0.6095, 0.6002, 0.6029, 0.6049, 0.6076, 0.6038, 0.601, 0.6054, 0.6047, 0.6053, 0.6023, 0.5997, 0.6077, 0.6023, 0.6042, 0.6033, 0.6045, 0.6012, 0.6008, 0.6038, 0.6072, 0.6029, 0.6025, 0.6019, 0.604, 0.6029, 0.6023, 0.6042, 0.6055, 0.6045, 0.6013, 0.6015, 0.6009, 0.6015, 0.602, 0.6012, 0.6026, 0.6031, 0.6023, 0.6019, 0.6027, 0.6022, 0.6024, 0.6023, 0.6022, 0.6025, 0.6026, 0.6023, 0.6021, 0.6022, 0.6021, 0.6024, 0.6023]\nbleu3 = [0.2744, 0.2757, 0.2997, 0.3086, 0.3159, 0.3104, 0.3152, 0.3225, 0.3189, 0.3263, 0.3259, 0.3274, 0.325, 0.3306, 0.329, 0.3272, 0.3288, 0.3307, 0.3295, 0.3288, 0.3318, 0.3276, 0.3303, 0.3276, 0.326, 0.3334, 0.3317, 0.3274, 0.3292, 0.331, 0.3308, 0.3336, 0.3309, 0.3302, 0.3293, 0.3315, 0.3321, 0.3272, 0.3328, 0.333, 0.3338, 0.3344, 0.3338, 0.3332, 0.3357, 0.3357, 0.3337, 0.3347, 0.3338, 0.3341, 0.3342, 0.3338, 0.3343, 0.3344, 0.3343, 0.3348, 0.3346, 0.3345, 0.3344, 0.3346, 0.3345]\nacc3 = [0.2158, 0.2238, 0.2233, 0.2261, 0.2234, 0.2267, 0.2275, 0.224, 0.221, 0.2295, 0.2248, 0.228, 0.223, 0.2223, 0.2252, 0.2236, 0.2246, 0.2232, 0.2258, 0.2277, 0.2218, 0.2249, 0.2274, 0.2258, 0.2229, 0.2288, 0.2237, 0.2235, 0.2224, 0.2259, 0.2251, 0.2273, 0.2252, 0.2259, 0.2243, 0.2242, 0.2213, 0.2246, 0.2272, 0.2241, 0.2252, 0.2241, 0.226, 0.2259, 0.2257, 0.2228, 0.2242, 0.2276, 0.2268, 0.2267, 0.2265, 0.2268, 0.2261, 0.2253, 0.2261, 0.2263, 0.2262, 0.2254, 0.226, 0.2251, 0.2251]\navg_time = 261\n\n#ViT\ntrain_loss4 = [2.0639, 1.5379, 1.4074, 1.3104, 1.2358, 1.1799, 1.1336, 1.1011, 1.0708, 1.0438, 1.0199, 0.9962, 0.9747, 0.956, 0.9367, 0.919, 0.9044, 0.8899, 0.8735, 0.8601, 0.8466, 0.8366, 0.823, 0.8107, 0.8008, 0.7912, 0.783, 0.7721, 0.7622, 0.7522, 0.7443, 0.7337, 0.7271, 0.7192, 0.7118, 0.7044, 0.6964, 0.6254, 0.599, 0.5866, 0.5788, 0.5715, 0.5666, 0.5616, 0.5566, 0.5525, 0.5493, 0.5466, 0.5375, 0.5362, 0.535, 0.5342, 0.5329, 0.5328, 0.5327, 0.5318, 0.5320, 0.5312, 0.5301, 0.5297]\nval_loss4 = [1.6, 1.4146, 1.3149, 1.2222, 1.1625, 1.1171, 1.0898, 1.0668, 1.0478, 1.0379, 1.0134, 0.9952, 0.9814, 0.9762, 0.9611, 0.9538, 0.9522, 0.95, 0.9385, 0.939, 0.9326, 0.9357, 0.9322, 0.934, 0.9305, 0.9284, 0.9381, 0.9416, 0.9414, 0.9394, 0.945, 0.9404, 0.9494, 0.9568, 0.9564, 0.9622, 0.9635, 0.956, 0.9626, 0.9698, 0.9759, 0.9803, 0.9857, 0.9891, 0.9927, 0.9993, 1.0037, 1.0071, 1.0066, 1.0076, 1.0078, 1.0088, 1.0091, 1.0093, 1.0095, 1.0109, 1.0106, 1.0112, 1.0121, 1.0122]\nscore4 = [79.7525, 77.9945, 78.788, 77.311, 77.0895, 76.559, 76.694, 76.0615, 76.089, 76.7855, 76.234, 76.3495, 76.0345, 76.04, 76.064, 76.3325, 76.4745, 75.9995, 76.3995, 76.478, 76.3775, 76.2445, 76.1515, 76.3555, 75.8785, 76.0625, 75.8265, 76.1635, 76.1605, 75.948, 76.0375, 75.7795, 75.872, 75.868, 76.2415, 75.939, 75.9925, 76.038, 75.971, 75.9935, 76.004, 75.9605, 75.9455, 75.9765, 75.9335, 75.807, 75.8485, 75.831, 75.8615, 75.85, 75.826]\nscore_norm4 = [0.6323, 0.6167, 0.6231, 0.6114, 0.6091, 0.6056, 0.6062, 0.6022, 0.6015, 0.6075, 0.603, 0.6041, 0.6014, 0.6019, 0.6024, 0.6046, 0.6054, 0.6014, 0.6046, 0.605, 0.6047, 0.6038, 0.6027, 0.6045, 0.6008, 0.6026, 0.6003, 0.6025, 0.6027, 0.6015, 0.6023, 0.6002, 0.6007, 0.6008, 0.6038, 0.6015, 0.602, 0.6022, 0.6016, 0.6017, 0.602, 0.6016, 0.6016, 0.6017, 0.6013, 0.6004, 0.6008, 0.6006, 0.6008, 0.6007, 0.6006, 0.6005, 0.6007, 0.6005, 0.6007, 0.6007, 0.6006, 0.6004, 0.6007, 0.6006]\nbleu4 = [0.2835, 0.2898, 0.3007, 0.3037, 0.3147, 0.3119, 0.3119, 0.3186, 0.3225, 0.3225, 0.3218, 0.3253, 0.3265, 0.3222, 0.3235, 0.3266, 0.3287, 0.3317, 0.3246, 0.3283, 0.3246, 0.3242, 0.3348, 0.3297, 0.3294, 0.3294, 0.3316, 0.3351, 0.3299, 0.3297, 0.329, 0.3325, 0.3324, 0.3319, 0.3334, 0.3314, 0.3289, 0.3327, 0.3343, 0.3336, 0.3333, 0.3332, 0.3337, 0.3342, 0.3334, 0.3344, 0.3342, 0.3335, 0.3342, 0.3345, 0.3342, 0.3344, 0.3346, 0.3340, 0.3343, 0.3344, 0.3341, 0.3343, 0.3342, 0.3343]\nacc4 = [0.2148, 0.22, 0.2163, 0.2239, 0.221, 0.2242, 0.227, 0.2247, 0.2249, 0.223, 0.2269, 0.2246, 0.2293, 0.2278, 0.2231, 0.2233, 0.2223, 0.2281, 0.2236, 0.2213, 0.2234, 0.2238, 0.2274, 0.2229, 0.2261, 0.2253, 0.2247, 0.2268, 0.224, 0.2287, 0.2253, 0.2276, 0.2259, 0.2288, 0.2267, 0.227, 0.224, 0.2247, 0.2258, 0.226, 0.2271, 0.2254, 0.2279, 0.2265, 0.2273, 0.2274, 0.2257, 0.2255, 0.2268, 0.2268, 0.2277, 0.2277, 0.2277, 0.2277, 0.2273, 0.2276, 0.2274, 0.2272, 0.2268, 0.2272]\navg_time = 253\n\n#mobile\ntrain_loss5 = [2.1084, 1.5455, 1.4021, 1.3071, 1.2299, 1.1695, 1.1224, 1.0872, 1.0558, 1.025, 0.9975, 0.9707, 0.9473, 0.9233, 0.9044, 0.885, 0.8673, 0.8478, 0.8306, 0.8162, 0.8004, 0.786, 0.7711, 0.7632, 0.7493, 0.7367, 0.7273, 0.7149, 0.7035, 0.6937, 0.6815, 0.672, 0.6641, 0.6538, 0.5768, 0.5486, 0.5381, 0.5299, 0.5226, 0.5159, 0.5106, 0.5052, 0.5024, 0.4968, 0.4943, 0.4844, 0.4821, 0.481, 0.4801, 0.4799, 0.4786, 0.4767, 0.4777, 0.4764, 0.4759, 0.4756, 0.4743, 0.4746, 0.4751, 0.4733, 0.4744]\nval_loss5 = [1.6162, 1.4165, 1.3033, 1.2239, 1.1613, 1.0962, 1.0975, 1.0446, 1.0341, 1.0117, 0.9828, 0.9759, 0.9615, 0.9674, 0.9371, 0.9539, 0.9388, 0.9293, 0.9391, 0.9341, 0.9193, 0.9322, 0.9091, 0.9416, 0.9223, 0.93, 0.9274, 0.941, 0.9347, 0.931, 0.9352, 0.9522, 0.9488, 0.9451, 0.9502, 0.961, 0.9653, 0.9727, 0.9795, 0.9825, 0.9933, 0.9992, 1.0054, 1.0102, 1.0136, 1.0126, 1.0149, 1.0164, 1.0182, 1.0184, 1.018, 1.0204, 1.0212, 1.0226, 1.023, 1.0226, 1.0226, 1.0233, 1.0229, 1.0236, 1.0235]\nscore5 = [78.635, 78.8425, 79.8415, 79.4625, 78.6745, 76.818, 78.069, 77.1805, 78.2385, 77.788, 76.6465, 76.992, 76.2205, 77.349, 77.2625, 78.217, 78.2055, 77.2365, 78.315, 76.997, 76.39, 76.0655, 76.6865, 77.1305, 76.2355, 76.3675, 76.1015, 76.441, 76.2855, 76.1925, 76.0675, 75.945, 76.404, 76.507, 76.3235, 76.231, 76.155, 76.18, 76.2555, 76.2095, 76.34, 76.217, 76.179, 76.1845, 76.222, 76.204, 76.257, 76.2015, 76.225, 76.2415, 76.2375, 76.2075, 76.2385, 76.238, 76.2215, 76.2425, 76.278, 76.2715, 76.263, 76.2105, 76.2615]\nscore_norm5 = [0.6226, 0.6242, 0.6313, 0.6278, 0.6229, 0.6083, 0.618, 0.6102, 0.6192, 0.6159, 0.6063, 0.6089, 0.603, 0.6127, 0.611, 0.618, 0.6183, 0.6109, 0.6186, 0.6092, 0.6044, 0.6023, 0.6074, 0.6097, 0.6039, 0.6055, 0.6029, 0.6055, 0.6044, 0.6036, 0.6023, 0.6016, 0.6047, 0.6056, 0.6042, 0.6036, 0.603, 0.6032, 0.6037, 0.6032, 0.6045, 0.6035, 0.6033, 0.6031, 0.6035, 0.6034, 0.6038, 0.6034, 0.6036, 0.6037, 0.6037, 0.6034, 0.6037, 0.6036, 0.6035, 0.6037, 0.604, 0.604, 0.6039, 0.6035, 0.6039]\nbleu5 = [0.2632, 0.2853, 0.2944, 0.3059, 0.3069, 0.3086, 0.3165, 0.3173, 0.3198, 0.3174, 0.3236, 0.3238, 0.3279, 0.3234, 0.331, 0.3268, 0.3314, 0.3312, 0.3378, 0.3328, 0.3345, 0.333, 0.3318, 0.338, 0.3287, 0.3287, 0.3315, 0.3309, 0.3319, 0.3346, 0.3324, 0.3333, 0.3322, 0.3302, 0.3326, 0.3321, 0.3335, 0.3331, 0.3328, 0.3337, 0.3328, 0.3332, 0.3322, 0.3322, 0.3313, 0.3331, 0.3327, 0.3327, 0.3329, 0.333, 0.3329, 0.3327, 0.3329, 0.3329, 0.3329, 0.3326, 0.3325, 0.3324, 0.3325, 0.3327, 0.3327]\nacc5 = [0.2209, 0.2171, 0.2085, 0.2159, 0.2155, 0.2195, 0.2178, 0.2172, 0.2172, 0.2168, 0.2209, 0.2227, 0.2228, 0.2176, 0.2166, 0.2161, 0.2148, 0.2204, 0.2189, 0.2194, 0.2243, 0.2204, 0.2209, 0.2201, 0.2256, 0.2233, 0.2249, 0.2208, 0.2241, 0.2238, 0.2247, 0.2241, 0.2245, 0.2216, 0.224, 0.2237, 0.2236, 0.2234, 0.2255, 0.2241, 0.2237, 0.2226, 0.2258, 0.2258, 0.225, 0.2247, 0.226, 0.2254, 0.226, 0.2249, 0.2262, 0.2256, 0.2266, 0.2263, 0.2267, 0.2264, 0.226, 0.2263, 0.2259, 0.2263, 0.227]\navg_time = 159\n\n\n# resnet\ntrain_loss6 = [2.0915, 1.5325, 1.3952, 1.3102, 1.2461, 1.1899, 1.1464, 1.1111, 1.0783, 1.0486, 1.0221, 0.9971, 0.9753, 0.9555, 0.935, 0.9196, 0.9014, 0.8888, 0.8719, 0.8579, 0.8462, 0.8345, 0.8214, 0.8097, 0.7989, 0.7873, 0.777, 0.7666, 0.7555, 0.7462, 0.7365, 0.7249, 0.7132, 0.6382, 0.6106, 0.5963, 0.5893, 0.5816, 0.5762, 0.5704, 0.5653, 0.5615, 0.5571, 0.554, 0.5436, 0.5421, 0.5409, 0.5409, 0.5402, 0.5388, 0.539, 0.5378, 0.5378, 0.5379, 0.5353, 0.5345, 0.5347, 0.5342, 0.535, 0.5342, 0.5345]\nval_loss6 = [1.6073, 1.4071, 1.3063, 1.2426, 1.1821, 1.141, 1.1095, 1.0831, 1.0549, 1.0399, 1.0177, 1.0118, 0.9955, 0.9912, 0.977, 0.9686, 0.9666, 0.9625, 0.9721, 0.9556, 0.9527, 0.9463, 0.9468, 0.9531, 0.9524, 0.9547, 0.9568, 0.9595, 0.9548, 0.9646, 0.9707, 0.9668, 0.9746, 0.9703, 0.9739, 0.9843, 0.9894, 0.996, 0.9995, 1.0031, 1.0107, 1.0143, 1.0183, 1.0211, 1.0224, 1.0239, 1.0247, 1.0251, 1.0259, 1.0271, 1.0268, 1.0282, 1.0286, 1.0299, 1.0303, 1.0303, 1.0305, 1.0303, 1.0308, 1.0304, 1.0309]\nscore6 = [78.28, 77.6415, 77.2735, 76.704, 76.415, 76.35, 75.92, 75.7185, 75.9105, 76.168, 76.0975, 75.965, 75.815, 76.1675, 76.563, 75.984, 75.813, 76.657, 76.8525, 76.007, 76.3195, 76.3745, 76.1585, 76.8525, 76.401, 76.652, 75.945, 75.9005, 75.87, 76.3215, 76.375, 75.879, 75.878, 76.128, 76.1325, 75.975, 76.058, 76.124, 75.9295, 76.0225, 75.9675, 75.9305, 76.012, 76.001, 76.0635, 76.0525, 76.0075, 76.0195, 76.0355, 76.031, 76.021, 76.0075, 76.0445, 76.0275, 76.0755, 76.04, 76.058, 76.063, 76.031, 76.0605, 76.0465]\nscore_norm6 = [0.6199, 0.616, 0.612, 0.6077, 0.6048, 0.6045, 0.6011, 0.5998, 0.6009, 0.6043, 0.6026, 0.602, 0.5998, 0.6043, 0.6073, 0.6017, 0.6008, 0.6076, 0.6101, 0.6023, 0.6046, 0.6048, 0.6033, 0.6087, 0.6051, 0.6069, 0.6012, 0.6014, 0.6015, 0.6052, 0.6057, 0.6012, 0.6015, 0.6035, 0.6037, 0.6022, 0.6029, 0.6033, 0.6018, 0.6025, 0.6022, 0.6016, 0.6025, 0.6023, 0.6028, 0.6028, 0.6024, 0.6025, 0.6027, 0.6026, 0.6025, 0.6024, 0.6027, 0.6026, 0.603, 0.6027, 0.6028, 0.6029, 0.6026, 0.6029, 0.6028]\nbleu6 = [0.2763, 0.3029, 0.2951, 0.3059, 0.316, 0.3105, 0.3162, 0.3158, 0.3254, 0.3182, 0.3193, 0.3201, 0.3197, 0.3212, 0.3184, 0.3206, 0.3233, 0.3181, 0.3173, 0.3238, 0.3252, 0.3268, 0.3234, 0.32, 0.3262, 0.3283, 0.3266, 0.3285, 0.3282, 0.3259, 0.3221, 0.3272, 0.3263, 0.3263, 0.3275, 0.3281, 0.3274, 0.3275, 0.3287, 0.3294, 0.3279, 0.3292, 0.3285, 0.3294, 0.3285, 0.3282, 0.3289, 0.3289, 0.3284, 0.3286, 0.3285, 0.3288, 0.3286, 0.3285, 0.3286, 0.3286, 0.3285, 0.3285, 0.3285, 0.3284, 0.3286]\nacc6 = [0.2242, 0.2214, 0.2184, 0.2261, 0.2314, 0.228, 0.2293, 0.2304, 0.2278, 0.2192, 0.2315, 0.2242, 0.2287, 0.2226, 0.2177, 0.2239, 0.2231, 0.2179, 0.2179, 0.2197, 0.2232, 0.2212, 0.2216, 0.2182, 0.2232, 0.2193, 0.2253, 0.2248, 0.2256, 0.2198, 0.2194, 0.2229, 0.2227, 0.2243, 0.2242, 0.2243, 0.2233, 0.2239, 0.2245, 0.2236, 0.2239, 0.2247, 0.2239, 0.2238, 0.223, 0.2226, 0.2236, 0.2235, 0.2233, 0.2236, 0.2239, 0.2242, 0.2232, 0.2238, 0.2236, 0.2239, 0.2241, 0.2238, 0.224, 0.2238, 0.2243]\navg_time = 180\n\n\n# b5\ntrain_loss7 = [1.8896, 1.4661, 1.3497, 1.2639, 1.2014, 1.152, 1.1167, 1.0901, 1.0605, 1.0372, 1.0142, 0.9934, 0.976, 0.9606, 0.9468, 0.9297, 0.9188, 0.907, 0.8926, 0.8818, 0.8712]\nval_loss7 = [1.527, 1.354, 1.2579, 1.1818, 1.1402, 1.1003, 1.0887, 1.0521, 1.0268, 1.0073, 0.9997, 0.9794, 0.9767, 0.9585, 0.9499, 0.9427, 0.9405, 0.9383, 0.9381, 0.9292, 0.9316]\nscore7 = [77.6655, 78.0615, 77.1445, 76.0065, 75.373, 75.94, 75.235, 74.9435, 75.4645, 75.3225, 75.0115, 75.458, 74.9695, 74.9275, 75.1345, 74.7755, 75.6835, 74.7535, 75.029, 75.1355, 75.1585]\nscore_norm7 = [0.614, 0.6188, 0.6118, 0.6032, 0.5967, 0.6024, 0.5957, 0.5933, 0.5975, 0.5963, 0.5946, 0.5978, 0.5937, 0.5936, 0.5955, 0.5925, 0.6002, 0.5924, 0.5946, 0.5952, 0.5955]\nbleu7 = [0.2832, 0.2756, 0.3083, 0.3077, 0.3155, 0.3177, 0.3241, 0.3249, 0.3198, 0.3221, 0.3229, 0.3282, 0.3268, 0.3285, 0.3271, 0.3308, 0.325, 0.3353, 0.3291, 0.3352, 0.3282]\nacc7 = [0.2272, 0.2145, 0.223, 0.2212, 0.2277, 0.2226, 0.2301, 0.2315, 0.2255, 0.2316, 0.2286, 0.2283, 0.2311, 0.23, 0.2302, 0.2311, 0.2226, 0.2344, 0.2304, 0.2305, 0.2274]\navg_time = 325\n\n# b2\ntrain_loss8 = [2.0893, 1.516, 1.3841, 1.2916, 1.2102, 1.1395, 1.0685, 1.0095, 0.9577, 0.9151, 0.8732, 0.8365, 0.7989, 0.7678, 0.7401, 0.7141, 0.6904, 0.6658, 0.6404, 0.622, 0.5996, 0.5819, 0.5641, 0.5461, 0.5322, 0.5129, 0.4962, 0.4843, 0.4695, 0.4549, 0.4436, 0.4308, 0.4189, 0.4068, 0.3976, 0.3889, 0.3773, 0.3669, 0.3589, 0.3466, 0.3423, 0.3325, 0.3197, 0.3141, 0.3072, 0.3079, 0.303, 0.2958, 0.2917, 0.2849, 0.2123, 0.1816, 0.1685, 0.1597, 0.1526, 0.1481, 0.1432, 0.1382, 0.1349, 0.131, 0.1278]\nval_loss8 = [1.5728, 1.3854, 1.29, 1.1959, 1.1291, 1.0663, 1.0179, 0.9577, 0.9168, 0.8949, 0.8439, 0.8133, 0.7845, 0.7601, 0.7591, 0.7225, 0.7051, 0.6974, 0.6912, 0.6671, 0.6582, 0.6462, 0.637, 0.6413, 0.6176, 0.6075, 0.5976, 0.6203, 0.5877, 0.5942, 0.5748, 0.5789, 0.5782, 0.5816, 0.5809, 0.5774, 0.5698, 0.5782, 0.5646, 0.5847, 0.5711, 0.5735, 0.5679, 0.5783, 0.5888, 0.5759, 0.5791, 0.5839, 0.5996, 0.5887, 0.5502, 0.5534, 0.5582, 0.5594, 0.5659, 0.5709, 0.5747, 0.5797, 0.5839, 0.5877, 0.5943]\nscore8 = [79.059, 78.137, 78.427, 76.805, 77.018, 76.2415, 76.637, 76.221, 76.1365, 75.9885, 75.4565, 75.485, 75.523, 75.499, 75.3955, 75.012, 75.0745, 75.038, 75.161, 74.7495, 75.229, 75.1255, 74.901, 74.619, 75.042, 74.8365, 74.9125, 74.697, 74.6615, 74.8435, 74.742, 74.644, 74.549, 74.8445, 74.6805, 74.6895, 74.577, 74.7, 74.7385, 74.7085, 74.405, 74.5175, 74.4685, 74.592, 74.6475, 74.5885, 74.5865, 74.6465, 74.5105, 74.744, 74.4645, 74.467, 74.4475, 74.413, 74.4385, 74.447, 74.4435, 74.4005, 74.413, 74.4485, 74.446]\nscore_norm8 = [0.6246, 0.6169, 0.6206, 0.6072, 0.6097, 0.6029, 0.6064, 0.6033, 0.6025, 0.6022, 0.5981, 0.5977, 0.599, 0.5982, 0.5976, 0.5941, 0.595, 0.5945, 0.5954, 0.5923, 0.5964, 0.5954, 0.5937, 0.5913, 0.5948, 0.5931, 0.5936, 0.5918, 0.5916, 0.5933, 0.5923, 0.5916, 0.5907, 0.5931, 0.5919, 0.5921, 0.5911, 0.5919, 0.5923, 0.592, 0.5895, 0.5905, 0.5902, 0.5913, 0.5917, 0.5914, 0.591, 0.5914, 0.5904, 0.5925, 0.5901, 0.5901, 0.5899, 0.5896, 0.5898, 0.5899, 0.5898, 0.5895, 0.5896, 0.5899, 0.5898]\nbleu8 = [0.2641, 0.301, 0.2813, 0.3089, 0.3129, 0.3138, 0.3186, 0.3214, 0.3231, 0.3234, 0.3295, 0.3313, 0.3301, 0.3332, 0.3355, 0.3396, 0.3389, 0.3394, 0.3374, 0.3417, 0.3362, 0.3392, 0.3406, 0.3451, 0.3427, 0.345, 0.3431, 0.3455, 0.3452, 0.3455, 0.3454, 0.3465, 0.3478, 0.346, 0.3473, 0.3471, 0.3477, 0.3478, 0.3469, 0.3463, 0.3496, 0.348, 0.349, 0.3482, 0.3474, 0.3489, 0.3495, 0.3487, 0.3498, 0.3475, 0.3515, 0.3509, 0.3518, 0.3521, 0.3518, 0.3512, 0.3524, 0.352, 0.352, 0.3515, 0.3518]\nacc8 = [0.2175, 0.221, 0.2178, 0.2226, 0.2186, 0.2223, 0.222, 0.2194, 0.2239, 0.2239, 0.2251, 0.2252, 0.2245, 0.2256, 0.2273, 0.2278, 0.2253, 0.2327, 0.2292, 0.2272, 0.2264, 0.2271, 0.2285, 0.2323, 0.2283, 0.2329, 0.2305, 0.234, 0.2344, 0.2313, 0.2291, 0.2342, 0.2331, 0.2336, 0.2317, 0.2322, 0.236, 0.2291, 0.2319, 0.2338, 0.2346, 0.2324, 0.2347, 0.2338, 0.2336, 0.2339, 0.2311, 0.2333, 0.2301, 0.2334, 0.2363, 0.2376, 0.235, 0.2378, 0.2351, 0.236, 0.2386, 0.2395, 0.2369, 0.2364, 0.2385]\navg_time = 178\n\n# b2+\ntrain_loss9 = [2.0668, 1.4546, 1.2462, 1.0749, 0.9329, 0.8107, 0.7025, 0.6267, 0.5548, 0.5042, 0.4652, 0.4292, 0.3909, 0.3599, 0.3375, 0.3167, 0.2903, 0.2754, 0.2529, 0.2478, 0.2298, 0.2224, 0.2135, 0.1999, 0.1945, 0.1877, 0.1808, 0.1744, 0.165, 0.163, 0.1586, 0.1579, 0.1484, 0.1457, 0.1427, 0.1375, 0.1344, 0.1304, 0.1229, 0.1225, 0.122, 0.1215, 0.1127, 0.1127, 0.1105, 0.1109, 0.1069, 0.1072, 0.1029, 0.101, 0.1002, 0.0992, 0.0996, 0.0927, 0.0982, 0.0905, 0.0943, 0.0882, 0.0885, 0.0869, 0.0866]\nval_loss9 = [1.5557, 1.3208, 1.0874, 0.9475, 0.8809, 0.6884, 0.6129, 0.5413, 0.488, 0.4429, 0.4349, 0.3773, 0.7696, 0.3398, 0.3299, 0.3094, 0.2956, 0.2898, 0.2843, 0.2816, 0.2734, 0.2592, 0.2583, 0.2526, 0.2375, 0.2441, 0.2512, 0.2319, 0.231, 0.237, 0.2285, 0.2277, 0.225, 0.218, 0.2223, 0.2299, 0.2243, 0.2265, 0.2234, 0.2357, 0.2303, 0.2282, 0.2177, 0.2239, 0.2186, 0.2258, 0.2192, 0.2182, 0.2305, 0.2431, 0.2236, 0.2158, 0.2272, 0.2247, 0.2312, 0.2336, 0.226, 0.2303, 0.2312, 0.2276, 0.2349]\nscore9 = [78.049, 78.051, 76.925, 76.6045, 75.989, 74.918, 74.8285, 74.6295, 74.522, 74.4975, 74.5385, 74.28, 74.8775, 74.2335, 74.2975, 74.2075, 74.1345, 74.123, 74.1695, 74.059, 74.035, 74.0545, 74.083, 74.068, 73.9, 73.9995, 74.0485, 74.0305, 73.981, 73.9545, 73.988, 73.977, 73.9275, 73.933, 73.968, 74.025, 73.9335, 73.968, 73.9905, 73.9445, 73.981, 73.9735, 74.0075, 73.973, 73.964, 73.954, 73.946, 74.008, 73.956, 73.933, 73.927, 73.925, 73.9185, 73.8835, 73.9385, 73.9695, 73.961, 73.92, 73.907, 73.885, 73.9125]\nscore_norm9 = [0.6165, 0.6194, 0.6097, 0.6062, 0.6025, 0.5936, 0.5932, 0.5913, 0.5904, 0.5905, 0.5905, 0.5888, 0.5948, 0.5885, 0.589, 0.5881, 0.5876, 0.5874, 0.5878, 0.5871, 0.5869, 0.5869, 0.5872, 0.5872, 0.5858, 0.5866, 0.5871, 0.5868, 0.5865, 0.5863, 0.5865, 0.5866, 0.5861, 0.5861, 0.5862, 0.5871, 0.5861, 0.5864, 0.5867, 0.5863, 0.5866, 0.5865, 0.5866, 0.5864, 0.5864, 0.5864, 0.5863, 0.5867, 0.5864, 0.5861, 0.5862, 0.586, 0.5861, 0.5858, 0.5862, 0.5863, 0.5864, 0.586, 0.586, 0.5859, 0.5859]\nbleu9 = [0.274, 0.2872, 0.3037, 0.3207, 0.3226, 0.3421, 0.3447, 0.3496, 0.3532, 0.3538, 0.3544, 0.3585, 0.3459, 0.3607, 0.3603, 0.3615, 0.3623, 0.363, 0.363, 0.3639, 0.3634, 0.3643, 0.3642, 0.3641, 0.366, 0.3651, 0.3652, 0.3653, 0.3665, 0.3655, 0.3656, 0.3657, 0.3669, 0.3658, 0.3661, 0.3657, 0.3666, 0.3669, 0.3662, 0.3663, 0.3667, 0.3666, 0.3668, 0.3668, 0.3675, 0.367, 0.3667, 0.3667, 0.3677, 0.3674, 0.3674, 0.3675, 0.3671, 0.3675, 0.3668, 0.3668, 0.3673, 0.3678, 0.3675, 0.368, 0.3672]\nacc9 = [0.2239, 0.2151, 0.218, 0.2218, 0.224, 0.2322, 0.2338, 0.2356, 0.239, 0.2447, 0.2376, 0.2472, 0.2353, 0.2486, 0.2481, 0.251, 0.2546, 0.2523, 0.2523, 0.2546, 0.2552, 0.2558, 0.2543, 0.2537, 0.2555, 0.255, 0.2551, 0.2534, 0.2534, 0.2562, 0.2544, 0.2539, 0.2546, 0.2549, 0.2528, 0.2541, 0.2549, 0.2542, 0.2552, 0.2567, 0.2583, 0.2548, 0.256, 0.257, 0.2573, 0.2538, 0.2548, 0.2542, 0.2579, 0.2553, 0.2593, 0.2587, 0.256, 0.255, 0.2567, 0.2594, 0.2562, 0.2568, 0.257, 0.2593, 0.257]\navg_time = 176","metadata":{"execution":{"iopub.status.busy":"2022-10-26T17:07:18.402557Z","iopub.status.idle":"2022-10-26T17:07:18.403356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x = range(1, 61)\nplt.figure()\nplt.plot(x, val_loss1[:60])\nplt.plot(x, val_loss3[:60])\nplt.plot(x, val_loss4[:60])\nplt.plot(x, val_loss5[:60])\nplt.plot(x, val_loss6[:60])\nplt.plot(x, val_loss8[:60])\n\nplt.ylabel(\"Val Loss\")\nplt.xlabel(\"Epochs\")\nplt.legend(['EfficientNet B0', 'TNT', 'ViT', 'MobileNet', 'ResNet50', 'EfficientNet B2'])\nplt.savefig(\"fig1.jpg\", dpi=200)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-10-26T17:07:18.404539Z","iopub.status.idle":"2022-10-26T17:07:18.405294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"window = 5\nplt.plot(range(61), score_norm1, color='tab:blue', alpha=0.2)\nplt.plot(range(61), score_norm3, color='tab:green', alpha=0.2)\nplt.plot(range(60), score_norm4, color='tab:red', alpha=0.2)\nplt.plot(range(61), score_norm5, color='tab:purple', alpha=0.2)\nplt.plot(range(61), score_norm6, color='tab:orange', alpha=0.2)\nplt.plot(range(61), score_norm8, color='tab:cyan', alpha=0.2)\nscore_norm1_a = [np.mean(score_norm1[i:i+window]) for i in range(61)]\nscore_norm3_a = [np.mean(score_norm3[i:i+window]) for i in range(61)]\nscore_norm4_a = [np.mean(score_norm4[i:i+window]) for i in range(61)]\nscore_norm5_a = [np.mean(score_norm5[i:i+window]) for i in range(61)]\nscore_norm6_a = [np.mean(score_norm6[i:i+window]) for i in range(61)]\nscore_norm8_a = [np.mean(score_norm8[i:i+window]) for i in range(61)]\n\nplt.plot(range(61), score_norm1_a, color='tab:blue', label='EfficientNet-B0')\nplt.plot(range(61), score_norm8_a, color='tab:cyan', label='EfficientNet-B2')\n\nplt.plot(range(61), score_norm3_a, color='tab:green', label='TNT')\nplt.plot(range(61), score_norm4_a, color='tab:red', label='ViT')\nplt.plot(range(61), score_norm5_a, color='tab:purple', label='MobileNet')\nplt.plot(range(61), score_norm6_a, color='tab:orange', label='ResNet50')\n\n\nplt.legend()\nplt.xlabel('Epochs')\nplt.ylabel('Normalized Levenshtein Distance')\nplt.savefig('6.jpg')","metadata":{"execution":{"iopub.status.busy":"2022-10-26T17:07:18.406393Z","iopub.status.idle":"2022-10-26T17:07:18.407103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(range(61), acc1, color='tab:blue', alpha=0.2)\nplt.plot(range(61), acc3, color='tab:green', alpha=0.2)\nplt.plot(range(60), acc4, color='tab:red', alpha=0.2)\nplt.plot(range(61), acc5, color='tab:purple', alpha=0.2)\nplt.plot(range(61), acc6, color='tab:orange', alpha=0.2)\nplt.plot(range(61), acc8, color='tab:cyan', alpha=0.2)\n\nacc1_a = [np.mean(acc1[i:i+window]) for i in range(61)]\nacc3_a = [np.mean(acc3[i:i+window]) for i in range(61)]\nacc4_a = [np.mean(acc4[i:i+window]) for i in range(61)]\nacc5_a = [np.mean(acc5[i:i+window]) for i in range(61)]\nacc6_a = [np.mean(acc6[i:i+window]) for i in range(61)]\nacc8_a = [np.mean(acc8[i:i+window]) for i in range(61)]\n\nplt.plot(range(61), acc1_a, color='tab:blue', label='EfficientNet-B0')\nplt.plot(range(61), acc8_a, color='tab:cyan', label='EfficientNet-B2')\n\nplt.plot(range(61), acc3_a, color='tab:green', label='TNT')\nplt.plot(range(61), acc4_a, color='tab:red', label='ViT')\nplt.plot(range(61), acc5_a, color='tab:purple', label='MobileNet')\nplt.plot(range(61), acc6_a, color='tab:orange', label='ResNet50')\n\n\nplt.legend()\nplt.xlabel('Epochs')\nplt.ylabel('Accuracy')\nplt.savefig('8.jpg')","metadata":{"execution":{"iopub.status.busy":"2022-10-26T17:07:18.408193Z","iopub.status.idle":"2022-10-26T17:07:18.408922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# B2 + Transformer on full dataset\nTrain_ACC= [0.84171337, 0.95416826, 0.9685068, 0.974872, 0.97877336, 0.9814853, 0.9837553, 0.9852735, 0.98689914, 0.98799366, 0.9888958, 0.9899165, 0.9906678, 0.99135846, 0.99189514, 0.9924564, 0.9928715, 0.9933255, 0.99366647, 0.99398047, 0.99427515, 0.9946619, 0.99480164, 0.9951643, 0.99525446, 0.99553174, 0.99573183, 0.99585336, 0.99601173, 0.99615353, 0.9962607, 0.99636585, 0.9964611, 0.99655676, 0.99660283, 0.99668485, 0.9967369, 0.9968021, 0.99689674, 0.9970108, 0.9970535, 0.99713176, 0.997206, 0.99724334, 0.9971818]\nTrain_LOSS= [0.31586358, 0.08660097, 0.060065743, 0.04762288, 0.039944395, 0.03473899, 0.030473031, 0.02773981, 0.024950046, 0.022967454, 0.021352723, 0.019624121, 0.018358637, 0.017138932, 0.016190914, 0.015240533, 0.014464627, 0.013661291, 0.013003343, 0.0123876035, 0.011873301, 0.011175033, 0.010869061, 0.010223911, 0.009987555, 0.009445464, 0.009086044, 0.008847369, 0.008515231, 0.008246343, 0.008005503, 0.007781884, 0.0075736465, 0.007425133, 0.0072889677, 0.007098362, 0.007004458, 0.00688394, 0.0066630924, 0.006468861, 0.006337913, 0.0061528734, 0.0060065337, 0.005933904, 0.006072221]\nVal_ACC= [0.6314258, 0.8494158, 0.83471656, 0.8809252, 0.90517014, 0.9127166, 0.91672564, 0.9145588, 0.93334997, 0.92979485, 0.9422037, 0.94641095, 0.9453356, 0.94709647, 0.94379693, 0.95329237, 0.95312005, 0.9543808, 0.95650464, 0.9553435, 0.96017766, 0.9592812, 0.9616285, 0.9614489, 0.96148777, 0.9623818, 0.9649265, 0.96105295, 0.9659457, 0.96723384, 0.9635515, 0.9666284, 0.9666594, 0.96821404, 0.9675132, 0.9672578, 0.96794415, 0.9673649, 0.96658033, 0.9691993, 0.96807384, 0.96815336, 0.9679187, 0.968482, 0.96895444]\nVal_LOSS= [2.9015675, 1.1803999, 1.2685184, 0.982094, 0.8181792, 0.76470345, 0.74052507, 0.7758008, 0.6384557, 0.64324695, 0.5675466, 0.53339106, 0.5481487, 0.5335885, 0.559286, 0.48587692, 0.48588014, 0.4841556, 0.4684765, 0.47654265, 0.43651247, 0.452352, 0.42442152, 0.43490005, 0.43789095, 0.43833405, 0.40574226, 0.4522567, 0.4021659, 0.38688585, 0.42927694, 0.40588307, 0.4022432, 0.391373, 0.4008707, 0.4031866, 0.39660332, 0.40621948, 0.42042542, 0.3902028, 0.40816975, 0.40810648, 0.4075573, 0.4061545, 0.40353936]\nLSD= [20.569462, 9.144819, 10.846905, 7.3615036, 5.8258715, 5.42046, 5.198405, 5.3606772, 4.012758, 4.4951797, 3.477514, 3.2378557, 3.3062525, 3.2733748, 3.5476136, 2.841947, 2.8644707, 2.811949, 2.6270032, 2.784455, 2.4434845, 2.4781024, 2.3332958, 2.3569586, 2.3513372, 2.2789838, 2.1228967, 2.409455, 2.061173, 1.9845127, 2.2348132, 2.0220103, 2.0435448, 1.9475536, 1.983724, 2.0063477, 1.9387771, 1.9788787, 2.0357697, 1.8845403, 1.946239, 1.9298252, 1.9707031, 1.9127854, 1.8992138]\nLSD_norm= [0.14905411, 0.066266805, 0.078600764, 0.053344242, 0.042216454, 0.03927869, 0.037669607, 0.038845498, 0.029077945, 0.032573763, 0.025199369, 0.023462728, 0.023958348, 0.023720106, 0.025707345, 0.020593816, 0.020757033, 0.020376438, 0.019036252, 0.020177215, 0.01770641, 0.017957266, 0.016907943, 0.017079413, 0.017038673, 0.01651438, 0.015383305, 0.017459823, 0.014936035, 0.014380524, 0.016194297, 0.014652248, 0.014808295, 0.014112711, 0.014374809, 0.014538753, 0.014049108, 0.014339698, 0.014751952, 0.013656083, 0.014103182, 0.013984239, 0.014280458, 0.013860764, 0.013762419]","metadata":{"execution":{"iopub.status.busy":"2022-10-26T17:07:18.410028Z","iopub.status.idle":"2022-10-26T17:07:18.410767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(range(45), Train_ACC, label=\"Training\", color='tab:blue')\nplt.plot(range(45), Val_ACC, label=\"Validation\", color='tab:orange')\n\nplt.legend()\nplt.xlabel('Epochs')\nplt.ylabel('Accuracy')\n# plt.savefig('9.jpg')\n","metadata":{"execution":{"iopub.status.busy":"2022-10-26T17:07:18.411845Z","iopub.status.idle":"2022-10-26T17:07:18.412567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax1 = plt.subplots()\n\nax2 = ax1.twinx()\nline1, = ax1.plot(range(45), LSD, 'tab:green', label='Levenshtein Distance')\nline2, = ax2.plot(range(45), LSD_norm, 'tab:red', label='Normalized Levenshtein Distance')\n\nax1.set_xlabel('Epochs')\nax1.set_ylabel('Levenshtein Distance')\nax1.set_ylim([0,25])\nax1.legend()\n\nax2.set_ylabel('Normalized Levenshtein Distance')\nax2.set_ylim([0,0.5])\nax1.legend(handles=[line1, line2])\n# plt.savefig('10.jpg')","metadata":{"execution":{"iopub.status.busy":"2022-10-26T17:07:18.413637Z","iopub.status.idle":"2022-10-26T17:07:18.414363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission (do not run the cells below)","metadata":{}},{"cell_type":"code","source":"# def inference(test_loader, encoder, decoder, tokenizer, device):\n    \n#     encoder.eval()\n#     decoder.eval()\n    \n#     text_preds = []\n#     tk0 = tqdm(test_loader, total = len(test_loader))\n    \n#     for images in tk0:\n        \n#         images = images.to(device)\n        \n#         with torch.no_grad():\n#             features = encoder(images)\n#             predictions = decoder.predict(features, CFG.max_len, tokenizer)\n            \n#         predicted_sequence = torch.argmax(predictions.detach().cpu(), -1).numpy()\n#         _text_preds = tokenizer.predict_captions(predicted_sequence)\n#         text_preds.append(_text_preds)\n        \n#     text_preds = np.concatenate(text_preds)\n    \n#     return text_preds","metadata":{"execution":{"iopub.status.busy":"2022-10-26T17:07:18.415423Z","iopub.status.idle":"2022-10-26T17:07:18.416135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# inference\n# ====================================================\n\n# test_dataset = TestDataset(test, transform = get_transforms(data = 'valid'))\n# test_loader  = DataLoader(test_dataset, batch_size = 256, shuffle = False, num_workers = CFG.num_workers)\n# predictions  = inference(test_loader, encoder, decoder, tokenizer, device)","metadata":{"execution":{"iopub.status.busy":"2022-10-26T17:07:18.417231Z","iopub.status.idle":"2022-10-26T17:07:18.417937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n#  submission\n# ====================================================\n\n# test['InChI'] = [f\"InChI=1S/{text}\" for text in predictions]\n# test[['image_id', 'InChI']].to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-10-26T17:07:18.418995Z","iopub.status.idle":"2022-10-26T17:07:18.419718Z"},"trusted":true},"execution_count":null,"outputs":[]}]}