{"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":"# Config","metadata":{}},{"cell_type":"code","source":"def showVersion():\n    import sys\n    import torch\n    import numpy as np\n    import pandas as pd\n    import cv2\n    import matplotlib\n    import albumentations\n    import timm\n    import sklearn\n    import tqdm\n\n    print('python: ' + sys.version)\n    print('torch: ' + torch.__version__)\n    print('numpy: ' + np.__version__)\n    print('pandas: ' + pd.__version__)\n    print('cv2: ' + cv2.__version__)\n    print('matplotlib: ' + matplotlib.__version__)\n    print('albumentations: ' + albumentations.__version__)\n    print('timm: ' + timm.__version__)\n    print('sklearn: ' + sklearn.__version__)\n    print('tqdm: ' + tqdm.__version__)\n    \n\n# showVersion()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n# !pip install einops","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:16:42.834164Z","iopub.execute_input":"2023-10-23T08:16:42.834821Z","iopub.status.idle":"2023-10-23T08:16:42.840922Z","shell.execute_reply.started":"2023-10-23T08:16:42.834744Z","shell.execute_reply":"2023-10-23T08:16:42.839462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PathConfig:\n    # 基本文件路径\n    BASE_DIR = '/kaggle'\n    INPUT_DIR = BASE_DIR + '/input'\n    WORKING_DIR = BASE_DIR + '/working'\n    \n    MT_DIR = INPUT_DIR + '/bms-molecular-translation'\n    PREPROCESSED_DIR = INPUT_DIR + '/preprocessed'\n    LOG_DIR = WORKING_DIR + '/log'\n    OUTPUT_DIR = WORKING_DIR + '/output'\n    \n    \n    # 原始数据集文件路径\n    TRAIN_DIR = MT_DIR + '/train'\n    TEST_DIR = MT_DIR + '/test'\n    TRAIN_CSV = MT_DIR + '/train_labels.csv'\n    TEST_CSV = MT_DIR + '/sample_submission.csv'\n    TEST_ORIENTATION_CSV = MT_DIR + '/test_orientation.csv'\n    \n\n    # 数据预处理得到的文件路径\n    TOKEN_STOI_PICKLE = PREPROCESSED_DIR + '/tokenizer.stoi.pickle'\n\n    TRAIN_PREPROCESSED_CSV = PREPROCESSED_DIR + '/train_preprocessed.csv'\n    VALID_PREPROCESSED_CSV = PREPROCESSED_DIR + '/valid_preprocessed.csv'\n    TEST_PREPROCESSED_CSV = PREPROCESSED_DIR + '/test_preprocessed.csv'\n    TRAIN_PREPROCESSED_PICKLE = PREPROCESSED_DIR + '/train_preprocessed.pickle'\n    VALID_PREPROCESSED_PICKLE = PREPROCESSED_DIR + '/valid_preprocessed.pickle'\n    TEST_PREPROCESSED_PICKLE = PREPROCESSED_DIR + '/test_preprocessed.pickle'\n\n    LOAD_WEIGHT_PATH = '/kaggle/input/vit-small-adam-epoch21/output/last_epoch.pth'\n    LOAD_OUTPUT = '/kaggle/input/vit-small-adam-epoch21/output/lr0.0001_batch64_encoderdim768_dropout0.5.csv'\n    # 输出文件路径\n    LAST_WEIGHT_PATH = OUTPUT_DIR + '/last_epoch.pth'\n    BEST_WEIGHT_PATH = OUTPUT_DIR + '/best.pth'\n\n\n# 训练时常量\nclass TrainConfig:\n    # 模型参数\n    ENCODER_DIM = 768  # 可换512\n    ENCODER_N_LAYER = 12\n    EMBED_DIM = 256\n    DECODER_DIM = 512\n    SIZE = 224\n\n    N_FOLD = 5\n    SEED = 42\n\n    BATCH_SIZE = 64\n    NUM_WORKERS = 2\n    EPOCHS = 24\n    LR = 1e-4\n    SCHEDULER_NAME = 'CosineAnnealingWarmRestarts'  # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts']\n    DROPOUT = 0.5\n    WEIGHT_DECAY = 1e-6\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    PRINT_FREQ = 1000\n    MAX_LEN = 275\n\n    START_EPOCH = 24\n\n\nclass Config:\n    PATH = PathConfig\n    TRAIN = TrainConfig\n","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:16:42.844004Z","iopub.execute_input":"2023-10-23T08:16:42.844786Z","iopub.status.idle":"2023-10-23T08:16:42.861463Z","shell.execute_reply.started":"2023-10-23T08:16:42.84474Z","shell.execute_reply":"2023-10-23T08:16:42.860106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import shutil\n# shutil.rmtree(Config.PATH.PREPROCESSED_DIR)","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:16:42.863721Z","iopub.execute_input":"2023-10-23T08:16:42.864166Z","iopub.status.idle":"2023-10-23T08:16:42.879637Z","shell.execute_reply.started":"2023-10-23T08:16:42.864125Z","shell.execute_reply":"2023-10-23T08:16:42.878188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\ndef mkdir(dir_path):\n    is_exists = os.path.exists(dir_path)\n    \n    if not is_exists: \n        os.makedirs(dir_path)","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:16:42.881422Z","iopub.execute_input":"2023-10-23T08:16:42.881934Z","iopub.status.idle":"2023-10-23T08:16:42.895318Z","shell.execute_reply.started":"2023-10-23T08:16:42.881889Z","shell.execute_reply":"2023-10-23T08:16:42.894257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mkdir(Config.PATH.LOG_DIR)\nmkdir(Config.PATH.OUTPUT_DIR)","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:16:42.899053Z","iopub.execute_input":"2023-10-23T08:16:42.900336Z","iopub.status.idle":"2023-10-23T08:16:42.914099Z","shell.execute_reply.started":"2023-10-23T08:16:42.900288Z","shell.execute_reply":"2023-10-23T08:16:42.912494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.path.exists(Config.PATH.TRAIN_PREPROCESSED_PICKLE)","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:16:42.916156Z","iopub.execute_input":"2023-10-23T08:16:42.917008Z","iopub.status.idle":"2023-10-23T08:16:42.929935Z","shell.execute_reply.started":"2023-10-23T08:16:42.916963Z","shell.execute_reply":"2023-10-23T08:16:42.928573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# utils","metadata":{}},{"cell_type":"code","source":"import math\nimport pickle\nimport time\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset\nimport Levenshtein\nfrom matplotlib import pyplot as plt\n# albumentations数据增强库\nfrom albumentations.pytorch import ToTensorV2\nimport albumentations\nfrom datetime import datetime\n# from src.utils.config import Config\n\nimport os\nimport random","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:16:42.932233Z","iopub.execute_input":"2023-10-23T08:16:42.933392Z","iopub.status.idle":"2023-10-23T08:16:42.942299Z","shell.execute_reply.started":"2023-10-23T08:16:42.933346Z","shell.execute_reply":"2023-10-23T08:16:42.94008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_torch(seed=42):\n    \"\"\"\n    主要用来设置各种包的随机种子, 将所有包的随机种子固定为一个值, 可以使得结果可复现\n    :param seed: \n    \"\"\"\n    random.seed(seed)\n    # os.environ 获取环境变量\n    # PYTHONHASHSEED python的hash种子\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # deterministic置为True的话，每次返回的卷积算法将是确定的，即默认算法。如果配合上设置 Torch 的随机种子为固定值的话，应该可以保证每次运行网络的时候相同输入的输出是固定的\n    torch.backends.cudnn.deterministic = True\n\n\nclass Tokenizer():\n    def __init__(self, filepath=None):\n        \"\"\"初始化方法\n        生成 stoi 字典和 itos 字典\n        stoi : {char:int} 字符和index映射\n        itos : {int:char} index和字符映射\n        \"\"\"\n        self.stoi = {}\n        self.itos = {}\n\n        if filepath:\n            with open(filepath, 'rb') as f:\n                self.stoi = pickle.load(f)\n            self.itos = {k: v for v, k in self.stoi.items()}\n\n    def __len__(self):\n        return len(self.stoi)\n\n    def create_dicts_for_texts(self, texts):\n        \"\"\"根据文本生成字典\n        :param: list, text 为分词后形成的 list e.g ['C 13 H 20 O S','C 21 H 30 O 4',...]\n        \"\"\"\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_seq(self, text):\n        \"\"\"将 text 转换成 int list, 加头<sos>尾<eos>\n            输入text='C 13 H 20 O S', 返回sequence=[190,98,0,23,4,54,43,191]\n        \"\"\"\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_seqs(self, texts):\n        \"\"\"将多个 text 转换成 intlist\n        \"\"\"\n        sequences = []\n        for text in texts:\n            sequence = self.text_to_seq(text)\n            sequences.append(sequence)\n        return sequences\n\n    def seq_to_text(self, sequence):\n        \"\"\"将 intlist 转换成 text\n            输入sequence=[190,98,0,23,4,54,43,191], 返回text='C 13 H 20 O S'\n        \"\"\"\n        return ''.join(list(map(lambda i: self.itos[i], sequence)))\n\n    def seqs_to_texts(self, sequences):\n        \"\"\"将多个 intlist 转换成text\n        \"\"\"\n        texts = []\n        for sequence in sequences:\n            text = self.seq_to_text(sequence)\n            texts.append(text)\n        return texts\n\n    def predict_caption(self, sequence):\n        \"\"\"将预测结果 (intlist) 转换为字符 (str)，组装为标准 InChI 格式\n        e.g [190, 178, 47, 182, 89, 185, 187, 6, 13, 4, 165, 0, 88, 1, 154, 4, 69, 4, 47, 4, 132, 4, 121, 4, 14, 0, 99, 1, 143, 4, 36, 0, 47, 1, 25, 0, 110, 1, 58, 7, 121, 4, 143, 3, 165, 3, 25, 3, 58, 182, 3, 154, 182, 88, 3, 13, 4, 110, 182, 99, 191]\n            ->\n            InChI=1S/C13H20OS/c1-9(2)8-15-13-6-5-10(3)7-12(13)11(4)14/h5-7,9,11,14H,8H2,1-4H3 \n        \"\"\"\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        \"\"\"将多个预测结果 (intlist) 转换为字符 (text)，组装为标准 InChI 格\n        \"\"\"\n        captions = []\n        for sequence in sequences:\n            caption = self.predict_caption(sequence)\n            captions.append(caption)\n        return captions\n\n    def get_seq_of_sos(self):\n        return self.stoi['<sos>']\n\n    def get_seq_of_eos(self):\n        return self.stoi['<eos>']\n\n    def get_seq_of_pad(self):\n        return self.stoi['<pad>']\n\n\nclass TrainDataset(Dataset):\n    def __init__(self, data_df: pd.DataFrame, filepath: str, transform):\n        super().__init__()\n        self.data_df = data_df\n        self.filepath = filepath\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data_df)\n\n    def __getitem__(self, index):\n        img_path = self.data_df['img_path'][index]\n        img_path = self.filepath + img_path\n        # 这里是以三通道的方式去读img的, 所以image直接就是三通道\n        # todo 直接读三通道和读一通道重复三次, 值是不一样的, 看后期是否要改\n        image = cv2.imread(img_path)\n\n        # 将BGR格式转换成RGB格式\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n        augmented = self.transform(image=image)\n        image_tsr = augmented['image']\n\n        label = self.data_df['seq'][index]\n        label_len = self.data_df['seq_len'][index]\n\n        # transform已经将image变为tensor了\n        return image_tsr, \\\n               torch.tensor(label).long(), \\\n               torch.tensor(label_len).long(), \\\n               self.data_df['InChI'][index]\n\n\nclass TestDataset(Dataset):\n    def __init__(self, data_df: pd.DataFrame, filepath: str, transform):\n        super().__init__()\n        self.data_df = data_df\n        self.filepath = filepath\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data_df)\n\n    def __getitem__(self, idx):\n        img_path = self.filepath + self.data_df['img_path'][idx]\n        image = cv2.imread(img_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n\n        augmented = self.transform(image=image)\n        image_tsr = augmented['image']\n\n        return image_tsr\n\n\ndef get_transforms():\n    return albumentations.Compose([\n        albumentations.Resize(Config.TRAIN.SIZE, Config.TRAIN.SIZE),\n        # 为什么要使用这个数值?\n        albumentations.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225],\n        ),\n        ToTensorV2(),\n    ])\n\n\ndef show_img(test_df):\n    plt.figure(figsize=(20, 20))\n\n    for i in range(20):\n        image = cv2.imread(test_df.loc[i, 'img_path'])\n        plt.subplot(5, 4, i + 1)\n        plt.imshow(image)\n\n    plt.show()\n\n\ndef show_trans_img(test_df, transform):\n    plt.figure(figsize=(20, 20))\n\n    for i in range(20):\n        image = cv2.imread(test_df.loc[i, 'img_path'])\n        h, w, _ = image.shape\n        if h > w:\n            image = transform(image=image)['image']\n        plt.subplot(5, 4, i + 1)\n        plt.imshow(image)\n\n    plt.show()\n\n\ndef get_logger(log_filepath):\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_filepath)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n\n    with open(file=log_filepath, mode=\"a\") as f:\n        f.seek(0)\n        f.truncate()\n\n    return logger\n\n\ndef get_score(y_true, y_pred):\n    scores = []\n    for true, pred in zip(y_true, y_pred):\n        score = Levenshtein.distance(true, pred)\n        scores.append(score)\n    avg_score = np.mean(scores)\n    return avg_score\n\n\ndef get_lcs(y_true, y_pred):\n    lcss = []\n    for true, pred in zip(y_true, y_pred):\n        lcs = longestCommonSubsequence(true, pred)\n        lcss.append(lcs / len(true))\n    avg_lcs = np.mean(lcss)\n    return avg_lcs\n\n\ndef get_cer(y_true, y_pred):\n    cers = []\n    for true, pred in zip(y_true, y_pred):\n        score = Levenshtein.distance(true, pred)\n        cer = float(score / len(true))\n        cers.append(cer)\n    return np.mean(cers)\n\n\ndef longestCommonSubsequence(text1: str, text2: str) -> int:\n    m, n = len(text1), len(text2)\n    dp = [[0] * (n + 1) for _ in range(m + 1)]\n\n    for i in range(1, m + 1):\n        for j in range(1, n + 1):\n            if text1[i - 1] == text2[j - 1]:\n                dp[i][j] = dp[i - 1][j - 1] + 1\n            else:\n                dp[i][j] = max(dp[i - 1][j], dp[i][j - 1])\n\n    return dp[m][n]\n\n\ndef scoring(label_texts, label_text_preds):\n    score = get_score(label_texts, label_text_preds)\n    lcs = get_lcs(label_texts, label_text_preds)\n    cer = get_cer(label_texts, label_text_preds)\n\n    return score, lcs, cer\n\n\nclass AverageMeter():\n    \"\"\"记录总数和平均数的类\"\"\"\n\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 time_since(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 asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef seq_loss_calculate(loss_fn, seq_preds, seq_truth, seq_lens):\n    # 要去掉真实token的第一个sos\n    seq_preds = pack_padded_sequence(seq_preds, seq_lens, batch_first=True).data\n    seq_truth = pack_padded_sequence(seq_truth, seq_lens, batch_first=True).data\n\n    return loss_fn(seq_preds, seq_truth)\n\n\ndef get_now_str():\n    return datetime.now().strftime('%Y_%m_%d')\n","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:16:43.08194Z","iopub.execute_input":"2023-10-23T08:16:43.082813Z","iopub.status.idle":"2023-10-23T08:16:43.143144Z","shell.execute_reply.started":"2023-10-23T08:16:43.082764Z","shell.execute_reply":"2023-10-23T08:16:43.142166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# model","metadata":{}},{"cell_type":"code","source":"import math\nimport timm\nimport torch\nfrom torch import nn\nfrom timm.models.layers import to_2tuple","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:16:43.145553Z","iopub.execute_input":"2023-10-23T08:16:43.146345Z","iopub.status.idle":"2023-10-23T08:16:43.155158Z","shell.execute_reply.started":"2023-10-23T08:16:43.146297Z","shell.execute_reply":"2023-10-23T08:16:43.153685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_name = 'vit_small_patch16_224'\npretrained_path = '/root/.cache/huggingface/hub/models--timm--vit_small_patch16_224.augreg_in21k_ft_in1k/S_16-i21k-300ep-lr_0.001-aug_light1-wd_0.03-do_0.0-sd_0.0--imagenet2012-steps_20k-lr_0.03-res_224.npz'\n\n\n# vit有12个block, 可以使用6个\nclass Encoder(nn.Module):\n    def __init__(self, encoder_dim=256, encoder_n_layer=8):\n        super().__init__()\n\n        vit = timm.create_model(\n            model_name,\n            pretrained=True\n        )\n\n        # 去掉后面几层\n        vit_n_layer = len(vit.blocks)\n        for i in range(vit_n_layer):\n            if i >= encoder_n_layer:\n                vit.blocks[i] = nn.Identity()\n        vit.head = nn.Identity()\n        print(vit)\n\n        self.encoder = nn.Sequential(\n            vit,\n            nn.Linear(384, encoder_dim),\n            nn.LayerNorm(encoder_dim, eps=1e-06, elementwise_affine=True),\n            nn.Dropout(p=0.1)\n        )\n\n    def forward(self, x):\n        \"\"\"\n        :param x: shape=[b,c,h,w] \n        \"\"\"\n\n        # shape=[b,encoder_dim]\n        return self.encoder(x)\n\n\nclass Decoder(nn.Module):\n    def __init__(self, encoder_dim, embed_dim, decoder_dim, vocab_size, dropout=0.5):\n        super().__init__()\n\n        self.vocab_size = vocab_size\n\n        self.embedding = nn.Embedding(vocab_size, embed_dim)\n        self.lstm_cell = nn.LSTMCell(embed_dim + encoder_dim, decoder_dim, bias=True)\n        self.fc_gate = nn.Sequential(\n            # lstm上一次的结果h和attention_out相乘, 使得带注意力的图像特征拥有了上一次seq的信息\n            # h要和attention_out相乘, 所以要从decoder_dim变为encoder_dim\n            nn.Linear(decoder_dim, encoder_dim),\n            nn.Sigmoid()\n        )\n\n        self.dropout = nn.Dropout(dropout)\n        self.fc_pred = nn.Linear(decoder_dim, vocab_size)\n        # 用均匀分布初始化一些层\n        self.init_weights()\n\n    def init_weights(self):\n        self.embedding.weight.data.uniform_(-0.1, 0.1)\n        self.fc_pred.bias.data.fill_(0)\n        self.fc_pred.weight.data.uniform_(-0.1, 0.1)\n\n    def forward(self, h, c, attention_out, seq, seq_lens):\n        \"\"\"\n        :param h: shape=[b,decoder_dim]\n        :param c: shape=[b,decoder_dim]\n        :param attention_out: shape=[b,n_patch,encoder_dim]\n        :param seq: shape=[b, b_max_len] (b_max_len代表当前batch中seq的最大len)\n        :param seq_lens: shape=[b, 1]\n        \"\"\"\n\n        # seq.shape=[b,b_max_len,embed_dim]\n        seq = self.embedding(seq)\n\n        b = seq.size(0)\n        preds = torch.zeros(b, max(seq_lens), self.vocab_size).to(seq.device)\n\n        # decode_lengths = [94, 92, 83, 79]\n        # 有几个lstm_cell = 循环几次 = 序列有几个词 = 序列的长度 = range(max(decode_lengths)) \n        for t in range(max(seq_lens)):\n            # 循环内部每次都要挑出第t个词放到lstm_cell中\n            batch_size_t = sum([l > t for l in seq_lens])\n\n            gate = self.fc_gate(h[:batch_size_t])\n            attention_out_gated = gate * attention_out[:batch_size_t]\n            h, c = self.lstm_cell(\n                # seq[:batch_size_t, t, :].shape=[b,embed_dim]\n                # 拼接后.shape=[b,embed_dim+encoder_dim]\n                torch.cat([seq[:batch_size_t, t, :], attention_out_gated], dim=1),\n                (h[:batch_size_t], c[:batch_size_t])\n            )  # (batch_size_t, decoder_dim)\n\n            # 对h这个预测结果做dropout可以模拟中间某个单词预测错误的情况, 增加泛化能力\n            # pred.shape=[b,vocab_size]\n            pred = self.fc_pred(self.dropout(h))\n            preds[:batch_size_t, t, :] = pred\n\n        return preds\n\n    def predict(self, h, c, attention_out, seq_max_len, start_token, end_token, pad_token):\n        b = attention_out.size(0)\n        token = torch.ones(b, dtype=torch.long).to(attention_out.device) * start_token\n\n        preds = torch.zeros(b, seq_max_len, self.vocab_size).to(attention_out.device)\n        for t in range(seq_max_len):\n            token_embed = self.embedding(token)\n\n            # attention_out.shape=[b,encoder_dim]\n            gate = self.fc_gate(h)\n            attention_out_gated = gate * attention_out\n\n            h, c = self.lstm_cell(\n                torch.cat([token_embed, attention_out_gated], dim=1),\n                (h, c)\n            )\n\n            # todo 这里已经做了argmax了, 所以preds[:, t, :]=token即可, 为了避免和训练代码对不上, 这里暂时不修改\n            # pred.shape=[b,vocab_size]\n            pred = self.fc_pred(h)\n            preds[:, t, :] = pred\n\n            token = torch.argmax(pred, -1)\n\n            if ((token == end_token) | (token == pad_token)).all():\n                break\n\n        return preds\n\n\nclass Image2InChI(nn.Module):\n    def __init__(\n            self, encoder_dim, encoder_n_layer,\n            embed_dim, decoder_dim, vocab_size):\n        super().__init__()\n        self.encoder = Encoder(encoder_dim, encoder_n_layer)\n        self.decoder = Decoder(encoder_dim, embed_dim, decoder_dim, vocab_size)\n\n        self.init_h = nn.Linear(encoder_dim, decoder_dim)\n        self.init_c = nn.Linear(encoder_dim, decoder_dim)\n\n    def forward(self, img, seq, seq_lens):\n        \"\"\"\n        :param img: x.shape=[b,c,h,w] \n        \"\"\"\n\n        # encoder_out.shape=[b,encoder_dim]\n        encoder_out = self.encoder(img)\n\n        # h0.shape=[b,decoder_dim]\n        # c0.shape=[b,decoder_dim]\n        h0, c0 = self.init_h_c(encoder_out)\n\n        decoder_out = self.decoder(h0, c0, encoder_out, seq, seq_lens)\n\n        return decoder_out\n\n    def predict(self, img, seq_max_len, start_token, end_token, pad_token):\n        \"\"\"\n        :param img: x.shape=[b,c,h,w] \n        \"\"\"\n\n        # encoder_out.shape=[b,encoder_dim]\n        encoder_out = self.encoder(img)\n\n        # h0.shape=[b,decoder_dim]\n        # c0.shape=[b,decoder_dim]\n        h0, c0 = self.init_h_c(encoder_out)\n\n        decoder_out = self.decoder.predict(h0, c0, encoder_out, seq_max_len, start_token, end_token, pad_token)\n\n        return decoder_out\n\n    def init_h_c(self, encoder_out):\n        \"\"\"\n        :param x: x.shape=[b,n_patch,dim] \n        \"\"\"\n\n        # h0.shape=[b,decoder_dim]  \n        # c0.shape=[b,decoder_dim]\n        h0 = self.init_h(encoder_out)\n        c0 = self.init_c(encoder_out)\n\n        return h0, c0\n","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:16:43.158222Z","iopub.execute_input":"2023-10-23T08:16:43.158748Z","iopub.status.idle":"2023-10-23T08:17:08.647191Z","shell.execute_reply.started":"2023-10-23T08:16:43.158703Z","shell.execute_reply":"2023-10-23T08:17:08.646028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 数据预处理","metadata":{}},{"cell_type":"code","source":"import re\nimport pandas as pd\nfrom sklearn.model_selection import StratifiedKFold\nfrom tqdm.auto import tqdm\n# from utils.utils import Tokenizer\n# from utils.config import Config\n\n\ntqdm.pandas()","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:17:08.650089Z","iopub.execute_input":"2023-10-23T08:17:08.650791Z","iopub.status.idle":"2023-10-23T08:17:08.657771Z","shell.execute_reply.started":"2023-10-23T08:17:08.650751Z","shell.execute_reply":"2023-10-23T08:17:08.656547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def split_formula(formula):\n    \"\"\"化学式预处理\n        :param: str 化学式\n        :return: str 分词，用空格分开 e.g C13H20OS -> C 13 H 20 O S\n    \"\"\"\n    string = ''\n    # 正则表达式获取以一个大写字母开头，任意多个小写字母和数字结尾的组合。 e.g C13 Br\n    for i in re.findall(r\"[A-Z][^A-Z]*\", formula):\n        # 匹配其中的字母\n        elem = re.match(r\"\\D+\", i).group()\n        # 得到其中的数字\n        num = i.replace(elem, \"\")\n        # 用空格做连接\n        if num == \"\":\n            string += f\"{elem} \"\n        else:\n            string += f\"{elem} {str(num)} \"\n    # 去除末尾空格\n    return string.rstrip(' ')\n\n\ndef split_text(text):\n    \"\"\"原子连接预处理\n    :param: str 原子连接式\n    :return: str 分词，用空格分开 e.g c1-9(2)8-15-13-6-5-10(3)7-12(13)11(4)14 -> /c 1 - 9 ( 2 ) 8 - 15 - 13 - 6 - 5 - 10 ( 3 ) 7 - 12 ( 13 ) 11 ( 4 ) 14\n    \"\"\"\n    string = ''\n    for i in re.findall(r\"[a-z][^a-z]*\", text):\n        elem = i[0]\n        num = i.replace(elem, \"\").replace('/', \"\")\n        num_string = ''\n        for j in re.findall(r\"[0-9]+[^0-9]*\", num):\n            num_list = list(re.findall(r'\\d+', j))\n            assert len(num_list) == 1, f\"len(num_list) != 1\"\n            _num = num_list[0]\n            if j == _num:\n                num_string += f\"{_num} \"\n            else:\n                extra = j.replace(_num, \"\")\n                num_string += f\"{_num} {' '.join(list(extra))} \"\n        string += f\"/{elem} {num_string}\"\n    return string.rstrip(' ')\n\n\ndef split_text_list(text):\n    \"\"\"原子连接处理\n    :param list 多个原子连接式\n    \"\"\"\n    string = ''\n    for formula in text:\n        string += ' ' + split_text(formula)\n    return string.rstrip(' ')\n\n\ndef preprocess_train_df(train_df: pd.DataFrame, tokenizer: Tokenizer) -> pd.DataFrame:\n    # /0/0/0/000011a64c74.png\n    train_df['img_path'] = train_df['image_id'].progress_apply(\n        lambda image_id: f'/{image_id[0]}/{image_id[1]}/{image_id[2]}/{image_id}.png'\n    )\n\n    # InChI=1S\n    train_df['InChI_prefix'] = train_df['InChI'].progress_apply(lambda inchi: inchi.split('/')[0])\n\n    # C10H15N5S\n    train_df['formula'] = train_df['InChI'].progress_apply(lambda inchi: inchi.split('/')[1])\n\n    # 拆分化学分子式和原子连接式 \n    # C 10 H 15 N 5 S /c 1 - 7 - 6 - 8 ( 9 ( 11 ) 12 ) 14 - 10 ( 13 - 7 ) 15 - 2 - 4 - 16 - 5 - 3 - 15 /h 6 H , 2 - 5 H 2 , 1 H 3 , ( H 3 , 11 , 12 )\n    train_df['text'] = train_df['formula'].progress_apply(\n        lambda formula: split_formula(formula)) + train_df['InChI'].progress_apply(\n        lambda inchi: split_text_list(inchi.split('/')[2:]))\n\n    # text转为seq\n    # [190,98,0,23,4,54,...,43,191]\n    train_df['seq'] = train_df['text'].progress_apply(lambda text: tokenizer.text_to_seq(text))\n\n    # 不包含 <sos> <eos>, 所以-2\n    train_df['seq_len'] = train_df['seq'].progress_apply(lambda seq: len(seq) - 2)\n\n    return train_df\n\n\ndef split2train_and_valid(data_df: pd.DataFrame) -> (pd.DataFrame, pd.DataFrame):\n    folds = data_df.copy()\n\n    # StratifiedKFold()\n    #   n_splits：默认为3，表示将数据划分为多少份，即k折交叉验证中的k；\n    #   random_state：默认为None，表示随机数的种子，只有当shuffle设置为True的时候才会生效。\n    stratified_k_fold = StratifiedKFold(n_splits=Config.TRAIN.N_FOLD, shuffle=True, random_state=Config.TRAIN.SEED)\n\n    # Fold.split(数据集, 按类别分层): 返回拆分后数据集的索引值, 即训练集和测试集的索引train_index, val_index\n    # 这里的类别是inchi的长度, 如果长度不够, 可能会导致n_splits>分层后的个数, 会报错, 所以将n_splits设置小一点\n    split_index_gen = stratified_k_fold.split(folds, folds['seq_len'])\n    train_index, valid_index = next(split_index_gen)\n\n    train_df = folds.loc[train_index]\n    valid_df = folds.loc[valid_index]\n\n    train_df = train_df.reset_index(drop=True)\n    valid_df = valid_df.reset_index(drop=True)\n\n    return train_df, valid_df\n\n\ndef preprocess_test_df(test_df: pd.DataFrame) -> pd.DataFrame:\n    # /0/0/0/000011a64c74.png\n    test_df['img_path'] = test_df['image_id'].progress_apply(\n        lambda image_id: f'/{image_id[0]}/{image_id[2]}/{image_id[2]}/{image_id}.png'\n    )\n\n    test_df = test_df.drop(columns='InChI')\n\n    return test_df\n","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:17:08.659662Z","iopub.execute_input":"2023-10-23T08:17:08.660052Z","iopub.status.idle":"2023-10-23T08:17:08.683929Z","shell.execute_reply.started":"2023-10-23T08:17:08.66002Z","shell.execute_reply":"2023-10-23T08:17:08.682912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main_preprocess():\n    train_df = pd.read_csv(Config.PATH.TRAIN_CSV)\n    print(f'train_df.shape: {train_df.shape}')\n    test_df = pd.read_csv(Config.PATH.TEST_CSV)\n    print(f'test_df.shape: {test_df.shape}')\n\n    tokenizer = Tokenizer(Config.PATH.TOKEN_STOI_PICKLE)\n    train_df = preprocess_train_df(train_df, tokenizer)\n    test_df = preprocess_test_df(test_df)\n\n    train_df, valid_df = split2train_and_valid(train_df)\n    print(f'train_df.shape: {train_df.shape}')\n    print(f'valid_df.shape: {valid_df.shape}')\n\n    train_df.to_csv(Config.PATH.TRAIN_PREPROCESSED_CSV)\n    valid_df.to_csv(Config.PATH.VALID_PREPROCESSED_CSV)\n    test_df.to_csv(Config.PATH.TEST_PREPROCESSED_CSV)\n\n    train_df.to_pickle(Config.PATH.TRAIN_PREPROCESSED_PICKLE)\n    valid_df.to_pickle(Config.PATH.VALID_PREPROCESSED_PICKLE)\n    test_df.to_csv(Config.PATH.TEST_PREPROCESSED_PICKLE)\n","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:17:08.685261Z","iopub.execute_input":"2023-10-23T08:17:08.685952Z","iopub.status.idle":"2023-10-23T08:17:08.704096Z","shell.execute_reply.started":"2023-10-23T08:17:08.68592Z","shell.execute_reply":"2023-10-23T08:17:08.702752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(TrainConfig.DEVICE)","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:17:08.705976Z","iopub.execute_input":"2023-10-23T08:17:08.706847Z","iopub.status.idle":"2023-10-23T08:17:08.7206Z","shell.execute_reply.started":"2023-10-23T08:17:08.706809Z","shell.execute_reply":"2023-10-23T08:17:08.719211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# main_preprocess()","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:17:08.722435Z","iopub.execute_input":"2023-10-23T08:17:08.723781Z","iopub.status.idle":"2023-10-23T08:17:08.733531Z","shell.execute_reply.started":"2023-10-23T08:17:08.723736Z","shell.execute_reply":"2023-10-23T08:17:08.732163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 训练","metadata":{}},{"cell_type":"code","source":"import time\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom timm.optim import Lookahead\nfrom torch import nn\nfrom torch.optim import RAdam, Adam\nfrom torch.utils.data import DataLoader\nfrom torch.nn.utils.rnn import pad_sequence, pack_padded_sequence\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau, CosineAnnealingLR, CosineAnnealingWarmRestarts","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:17:08.735601Z","iopub.execute_input":"2023-10-23T08:17:08.736458Z","iopub.status.idle":"2023-10-23T08:17:08.776911Z","shell.execute_reply.started":"2023-10-23T08:17:08.736412Z","shell.execute_reply":"2023-10-23T08:17:08.775793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PARAM_PATHNAME = f'/lr{Config.TRAIN.LR}_batch{Config.TRAIN.BATCH_SIZE}_encoderdim{Config.TRAIN.ENCODER_DIM}_dropout{Config.TRAIN.DROPOUT}'\n\nLOG = get_logger(Config.PATH.LOG_DIR + PARAM_PATHNAME + '.log')\n\n\ndef bms_collate(batch, tokenizer, is_sort_by_len_desc=True):\n    \"\"\"对一个batch进行填充\"\"\"\n    imgs, labels, label_lens, label_texts = [], [], [], []\n    for row in batch:\n        imgs.append(row[0])\n        labels.append(row[1])\n        label_lens.append(row[2])\n        label_texts.append(row[3])\n\n    img = torch.stack(imgs)\n    seq = pad_sequence(\n        sequences=labels,\n        batch_first=True,\n        padding_value=tokenizer.stoi['<pad>']\n    )\n    seq_lens = torch.stack(label_lens)\n    seq_texts = label_texts\n\n    if is_sort_by_len_desc:\n        img, seq, seq_lens, seq_texts = sort_by_len_desc(img, seq, seq_lens, label_texts)\n    return img, seq, seq_lens, seq_texts\n\n\ndef sort_by_len_desc(img, seq, seq_lens, seq_texts):\n    \"\"\"\n    按照seq的len进行排序\n    \"\"\"\n\n    seq_lens, sorted_index = seq_lens.sort(dim=0, descending=True)\n    img = img[sorted_index]\n    seq = seq[sorted_index]\n\n    texts = [None] * len(sorted_index)\n    for i, j in enumerate(sorted_index.tolist()):\n        texts[j] = seq_texts[i]\n\n    return img, seq, seq_lens, texts\n\n\ndef get_data_loader(\n        train_df: pd.DataFrame,\n        valid_df: pd.DataFrame,\n        tokenizer: Tokenizer) -> (DataLoader, DataLoader):\n    transforms = get_transforms()\n\n    train_ds = TrainDataset(train_df, Config.PATH.TRAIN_DIR, transforms)\n    valid_ds = TrainDataset(valid_df, Config.PATH.TRAIN_DIR, transforms)\n\n    train_dl = DataLoader(\n        train_ds,\n        batch_size=Config.TRAIN.BATCH_SIZE,\n        shuffle=True,\n        num_workers=Config.TRAIN.NUM_WORKERS,\n        pin_memory=True,\n        drop_last=True,\n        collate_fn=lambda batch: bms_collate(batch, tokenizer, True)\n    )\n    valid_dl = DataLoader(\n        valid_ds,\n        batch_size=Config.TRAIN.BATCH_SIZE,\n        shuffle=False,\n        num_workers=Config.TRAIN.NUM_WORKERS,\n        pin_memory=True,\n        drop_last=False,\n        collate_fn=lambda batch: bms_collate(batch, tokenizer, False)\n    )\n\n    return train_dl, valid_dl\n","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:17:08.781289Z","iopub.execute_input":"2023-10-23T08:17:08.782056Z","iopub.status.idle":"2023-10-23T08:17:08.802682Z","shell.execute_reply.started":"2023-10-23T08:17:08.782009Z","shell.execute_reply":"2023-10-23T08:17:08.801433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def init_net(net, optimizer, scheduler):\n    net.to(Config.TRAIN.DEVICE)\n\n    LOAD_WEIGHT_PATH = Config.PATH.LOAD_WEIGHT_PATH\n    if Config.TRAIN.START_EPOCH > 1:\n        states = torch.load(LOAD_WEIGHT_PATH, map_location=torch.device(Config.TRAIN.DEVICE))\n\n        net.load_state_dict(states['net'])\n        optimizer.load_state_dict(states['optimizer'])\n        scheduler.load_state_dict(states['scheduler'])\n\n        LOG.info(f'初始epoch{Config.TRAIN.START_EPOCH}, 加载权重文件: {LOAD_WEIGHT_PATH}')\n        \n        output_load = pd.read_csv(Config.PATH.LOAD_OUTPUT)\n        print(output_load)\n\n\ndef init_data_device(img, seq, seq_lens):\n    device = Config.TRAIN.DEVICE\n\n    img = img.to(device)\n    seq = seq.to(device)\n    seq_lens = seq_lens.to(device)\n    return img, seq, seq_lens\n\n\ndef get_scheduler(optimizer):\n    scheduler_name = Config.TRAIN.SCHEDULER_NAME\n    if scheduler_name == 'ReduceLROnPlateau':\n        return ReduceLROnPlateau(\n            optimizer,\n            mode='min',  # 'min’模式检测metric是否不再减小，'max’模式检测metric是否不再增大\n            factor=0.2,  # 触发条件后lr*=factor\n            patience=4,  # 不再减小（或增大）的累计次数\n            verbose=True,  # 触发条件后print\n            eps=1e-8  # 如果新旧lr之间的差异小与1e-8，则忽略此次更新\n        )\n\n    if scheduler_name == 'CosineAnnealingLR':\n        return CosineAnnealingLR(\n            optimizer,\n            T_max=4,  # max_epoch=40次，那么设置T_max=5则会让学习率余弦周期性变化4次.\n            eta_min=1e-6,\n            last_epoch=-1\n        )\n\n    if scheduler_name == 'CosineAnnealingWarmRestarts':\n        return CosineAnnealingWarmRestarts(\n            optimizer,\n            T_0=4,  # 学习率第一次回到初始值的epoch位置\n            T_mult=1,  # 控制了学习率变化的速度\n            eta_min=1e-6,\n            last_epoch=-1\n        )\n\n    return None\n\n\ndef scheduler_step(scheduler, score=None):\n    if scheduler is None:\n        return\n\n    if isinstance(scheduler, ReduceLROnPlateau):\n        scheduler.step(score)\n        return\n\n    if isinstance(scheduler, CosineAnnealingLR):\n        scheduler.step()\n        return\n\n    if isinstance(scheduler, CosineAnnealingWarmRestarts):\n        scheduler.step()\n\n\ndef do_train(train_dl, net, loss_fn, optimizer):\n    losses = AverageMeter()\n    net.train()\n\n    for step, (img, seq, seq_lens, seq_text) in enumerate(train_dl):\n        b = img.size(0)\n        img, seq, seq_lens = init_data_device(img, seq, seq_lens)\n        # 因为数据预处理len-2, 忽略了开头和结尾的长度, 所以这里应该不用减1, 反而应该+1\n        seq_lens = (seq_lens + 1).tolist()\n\n        preds = net(img, seq, seq_lens)\n\n        # 注意排序问题\n        # 这里真实的seq要忽略前面的sos, 因为预测的seq没有sos\n        loss = seq_loss_calculate(loss_fn, preds, seq[:, 1:], seq_lens)\n        optimizer.zero_grad()\n        loss.backward()\n        nn.utils.clip_grad_norm_(net.parameters(), max_norm=5)\n        optimizer.step()\n\n        losses.update(loss.item(), b)\n\n        if step % Config.TRAIN.PRINT_FREQ == 0 or step == (len(train_dl) - 1):\n            LOG.info(\n                f'Train: [{step + 1}/{len(train_dl)}], '\n                f'CurLoss: {losses.val:.4f}, '\n                f'AvgLoss: {losses.avg:.4f}'\n            )\n\n    return losses.avg\n\n\ndef do_valid(valid_dl, net, loss_fn, tokenizer):\n    net.eval()\n\n    seq_texts = []\n    seq_text_preds = []\n    for step, (img, seq, seq_lens, seq_text) in enumerate(valid_dl):\n        img = img.to(Config.TRAIN.DEVICE)\n\n        with torch.no_grad():\n            preds = net.predict(\n                img,\n                Config.TRAIN.MAX_LEN,\n                tokenizer.get_seq_of_sos(),\n                tokenizer.get_seq_of_eos(),\n                tokenizer.get_seq_of_pad()\n            )\n            # todo 计算valid的loss\n\n        predicted_sequence = torch.argmax(preds.detach().cpu(), -1).numpy()\n        label_text_pred = tokenizer.predict_captions(predicted_sequence)\n        seq_text_preds.append(label_text_pred)\n        seq_texts.append(seq_text)\n\n        if step % Config.TRAIN.PRINT_FREQ == 0 or step == (len(valid_dl) - 1):\n            LOG.info(\n                f'Valid: [{step + 1}/{len(valid_dl)}] '\n            )\n\n    return np.concatenate(seq_texts), np.concatenate(seq_text_preds)\n\n\ndef train_and_valid(\n        train_dl, valid_dl, net,\n        optimizer, loss_fn, tokenizer, scheduler):\n    f_name = Config.PATH.OUTPUT_DIR + PARAM_PATHNAME + '.csv'\n    if Config.TRAIN.START_EPOCH == 1:\n        with open(file=f_name, mode=\"a\") as f:\n            f.seek(0)\n            f.truncate()\n            f.write('epoch,score,lcs,cer\\n')\n\n    start_epoch = Config.TRAIN.START_EPOCH\n    LOG.info(f\"start_epoch: {start_epoch}\")\n\n    best_score = np.inf\n\n    for epoch in range(Config.TRAIN.EPOCHS - start_epoch + 1):\n        epoch = epoch + start_epoch\n        LOG.info(f'Epoch: {epoch} / {Config.TRAIN.EPOCHS}')\n        start_time = time.time()\n\n        avg_loss = do_train(train_dl, net, loss_fn, optimizer)\n\n        label_texts, label_text_preds = do_valid(valid_dl, net, loss_fn, tokenizer)\n\n        label_text_preds = [f'InChI=1S/{text}' for text in label_text_preds]\n        LOG.info(f\"label_texts: {label_texts[:5]}\")\n        LOG.info(f\"label_text_preds: {label_text_preds[:5]}\")\n\n        score, lcs, cer = scoring(label_texts, label_text_preds)\n\n        with open(file=f_name, mode=\"a\") as f:\n            f.write(f\"{epoch},{score:.4f},{lcs:.4f},{cer:.4f}\\n\")\n\n        scheduler_step(scheduler, score)\n\n        elapsed_time = time.time() - start_time\n\n        LOG.info(\n            f'Epoch {epoch} - avg_train_loss: {avg_loss:.4f} - Score: {score:.4f} - lcs: {lcs:.4f} - cer: {cer:.4f} - time: {elapsed_time:.0f}s')\n\n        save_net(net, optimizer, scheduler, Config.PATH.LAST_WEIGHT_PATH)\n        if score < best_score:\n            best_score = score\n            save_net(net, optimizer, scheduler, Config.PATH.BEST_WEIGHT_PATH)\n            LOG.info(f'Epoch {epoch} - Save Best Score: {best_score:.4f} - lcs: {lcs:.4f} Model')\n\n\ndef save_net(net, optimizer, scheduler, filepath):\n    torch.save(\n        {\n            'net': net.state_dict(),\n            'optimizer': optimizer.state_dict(),\n            'scheduler': scheduler.state_dict()\n        },\n        filepath\n    )\n","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:17:08.804444Z","iopub.execute_input":"2023-10-23T08:17:08.805646Z","iopub.status.idle":"2023-10-23T08:17:08.830374Z","shell.execute_reply.started":"2023-10-23T08:17:08.805607Z","shell.execute_reply":"2023-10-23T08:17:08.828867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main_train():\n    seed_torch(Config.TRAIN.SEED)\n    \n    train_df = pd.read_pickle(Config.PATH.TRAIN_PREPROCESSED_PICKLE)\n    valid_df = pd.read_pickle(Config.PATH.VALID_PREPROCESSED_PICKLE)\n    valid_df = valid_df.head(20000)\n    tokenizer = Tokenizer(Config.PATH.TOKEN_STOI_PICKLE)\n\n    train_dl, valid_dl = get_data_loader(train_df, valid_df, tokenizer)\n\n    inputs, targets, lens, label_text = next(iter(train_dl))\n    LOG.info(inputs.shape)  # torch.Size([b, 3, 224, 224])\n    LOG.info(targets.shape)  # torch.Size([b, 184])\n    LOG.info(lens.shape)  # torch.Size([b, 1])\n\n    net = Image2InChI(\n        encoder_dim=Config.TRAIN.ENCODER_DIM,\n        encoder_n_layer=Config.TRAIN.ENCODER_N_LAYER,\n        embed_dim=Config.TRAIN.EMBED_DIM,\n        decoder_dim=Config.TRAIN.DECODER_DIM,\n        vocab_size=len(tokenizer)\n    )\n\n    optimizer = Adam(net.parameters(), lr=Config.TRAIN.LR, weight_decay=Config.TRAIN.WEIGHT_DECAY)\n#     optimizer = Lookahead(\n#         RAdam(\n#             filter(lambda p: p.requires_grad, net.parameters()),\n#             lr=0.001\n#         ),\n#         alpha=0.5,\n#         k=5\n#     )\n\n    loss_fn = nn.CrossEntropyLoss(ignore_index=tokenizer.stoi['<pad>'])\n    scheduler = get_scheduler(optimizer)\n\n    init_net(net, optimizer, scheduler)\n\n    train_and_valid(\n        train_dl, valid_dl, net,\n        optimizer, loss_fn, tokenizer, scheduler\n    )\n    ","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:17:08.831779Z","iopub.execute_input":"2023-10-23T08:17:08.832597Z","iopub.status.idle":"2023-10-23T08:17:08.850574Z","shell.execute_reply.started":"2023-10-23T08:17:08.832546Z","shell.execute_reply":"2023-10-23T08:17:08.849406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"main_train()","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:17:08.852433Z","iopub.execute_input":"2023-10-23T08:17:08.853417Z","iopub.status.idle":"2023-10-23T08:18:03.931177Z","shell.execute_reply.started":"2023-10-23T08:17:08.853371Z","shell.execute_reply":"2023-10-23T08:18:03.929727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_valid():\n    transforms = get_transforms()\n    tokenizer = Tokenizer(Config.PATH.TOKEN_STOI_PICKLE)\n\n    valid_df = pd.read_pickle(Config.PATH.VALID_PREPROCESSED_PICKLE)\n    valid_df = valid_df.head(640)\n    valid_ds = TrainDataset(valid_df, Config.PATH.TRAIN_DIR, transforms)\n    valid_dl = DataLoader(\n        valid_ds,\n        batch_size=Config.TRAIN.BATCH_SIZE,\n        shuffle=False,\n        num_workers=Config.TRAIN.NUM_WORKERS,\n        pin_memory=True,\n        drop_last=False,\n        collate_fn=lambda batch: bms_collate(batch, tokenizer, False)\n    )\n\n    encoder = Encoder(Config.TRAIN.ENCODER_MODEL_NAME, pretrained=True)\n    encoder_optimizer = Adam(encoder.parameters(), lr=Config.TRAIN.ENCODER_LR, weight_decay=Config.TRAIN.WEIGHT_DECAY,\n                             amsgrad=False)\n    decoder = DecoderWithAttention(\n        attention_dim=Config.TRAIN.ATTENTION_DIM,\n        embed_dim=Config.TRAIN.EMBED_DIM,\n        decoder_dim=Config.TRAIN.DECODER_DIM,\n        vocab_size=len(tokenizer),\n        dropout=Config.TRAIN.DROPOUT,\n        device=Config.TRAIN.DEVICE\n    )\n    decoder_optimizer = Adam(decoder.parameters(), lr=Config.TRAIN.DECODER_LR, weight_decay=Config.TRAIN.WEIGHT_DECAY,\n                             amsgrad=False)\n    loss_fn = nn.CrossEntropyLoss(ignore_index=tokenizer.stoi[\"<pad>\"])\n\n    encoder.to(Config.TRAIN.DEVICE)\n    decoder.to(Config.TRAIN.DEVICE)\n\n    filepath = '/kaggle/input/inchi-pan/output/resnet50_best.pth'\n    states = torch.load(filepath, map_location=torch.device(Config.TRAIN.DEVICE))\n    encoder.load_state_dict(states['encoder'])\n    decoder.load_state_dict(states['decoder'])\n    LOG.info(f'加载权重文件: {filepath}')\n\n    label_texts, label_text_preds = do_valid(valid_dl, encoder, decoder, tokenizer)\n\n    label_text_preds = [f'InChI=1S/{text}' for text in label_text_preds]\n    LOG.info(f\"label_texts: {label_texts[:5]}\")\n    LOG.info(f\"label_text_preds: {label_text_preds[:5]}\")\n\n    score, lcs, cer = scoring(label_texts, label_text_preds)\n\n    LOG.info(f'score: {score}, lcs: {lcs}, cer: {cer}')\n    \n\n# test_valid()","metadata":{"execution":{"iopub.status.busy":"2023-10-23T08:18:03.933295Z","iopub.execute_input":"2023-10-23T08:18:03.933774Z","iopub.status.idle":"2023-10-23T08:18:03.95146Z","shell.execute_reply.started":"2023-10-23T08:18:03.933732Z","shell.execute_reply":"2023-10-23T08:18:03.950319Z"},"trusted":true},"execution_count":null,"outputs":[]}]}