{"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"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":22422,"databundleVersionId":2153105,"sourceType":"competition"},{"sourceId":6681008,"sourceType":"datasetVersion","datasetId":3853819},{"sourceId":6896669,"sourceType":"datasetVersion","datasetId":3961713},{"sourceId":7087668,"sourceType":"datasetVersion","datasetId":3965421},{"sourceId":7019134,"sourceType":"datasetVersion","datasetId":3973702},{"sourceId":6924799,"sourceType":"datasetVersion","datasetId":3976099},{"sourceId":7033722,"sourceType":"datasetVersion","datasetId":3981376}],"dockerImageVersionId":30559,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"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":{"iopub.status.busy":"2023-11-07T00:45:03.547048Z","iopub.execute_input":"2023-11-07T00:45:03.547418Z","iopub.status.idle":"2023-11-07T00:45:03.560934Z","shell.execute_reply.started":"2023-11-07T00:45:03.547329Z","shell.execute_reply":"2023-11-07T00:45:03.559696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch","metadata":{"execution":{"iopub.status.busy":"2023-11-07T00:45:03.562598Z","iopub.execute_input":"2023-11-07T00:45:03.562965Z","iopub.status.idle":"2023-11-07T00:45:07.179832Z","shell.execute_reply.started":"2023-11-07T00:45:03.562935Z","shell.execute_reply":"2023-11-07T00:45:07.178811Z"},"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/swint-fairseq-adam/output/last_epoch.pth'\n    LOAD_OUTPUT = '/kaggle/input/swint-fairseq-adam/output/lr0.0001_batch32_encoderdim512_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 = 512\n    EMBED_DIM = 512\n    N_HEAD = 8\n    FF_DIM = 1024\n    NUM_LAYER = 3\n    SIZE = 224\n\n    N_FOLD = 5\n    SEED = 42\n\n    BATCH_SIZE = 32\n    NUM_WORKERS = 2\n    EPOCHS = 32\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 = 300\n\n    START_EPOCH = 32\n\n\nclass Config:\n    PATH = PathConfig\n    TRAIN = TrainConfig\n","metadata":{"execution":{"iopub.status.busy":"2023-11-07T00:45:07.181199Z","iopub.execute_input":"2023-11-07T00:45:07.181591Z","iopub.status.idle":"2023-11-07T00:45:07.214325Z","shell.execute_reply.started":"2023-11-07T00:45:07.181564Z","shell.execute_reply":"2023-11-07T00:45:07.213116Z"},"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-11-07T00:45:07.216968Z","iopub.execute_input":"2023-11-07T00:45:07.217656Z","iopub.status.idle":"2023-11-07T00:45:07.243888Z","shell.execute_reply.started":"2023-11-07T00:45:07.217618Z","shell.execute_reply":"2023-11-07T00:45:07.242859Z"},"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-11-07T00:45:07.245255Z","iopub.execute_input":"2023-11-07T00:45:07.245638Z","iopub.status.idle":"2023-11-07T00:45:07.255785Z","shell.execute_reply.started":"2023-11-07T00:45:07.245601Z","shell.execute_reply":"2023-11-07T00:45:07.254755Z"},"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-11-07T00:45:07.257583Z","iopub.execute_input":"2023-11-07T00:45:07.257949Z","iopub.status.idle":"2023-11-07T00:45:07.267793Z","shell.execute_reply.started":"2023-11-07T00:45:07.257916Z","shell.execute_reply":"2023-11-07T00:45:07.266629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.path.exists(Config.PATH.TRAIN_PREPROCESSED_PICKLE)","metadata":{"execution":{"iopub.status.busy":"2023-11-07T00:45:07.270934Z","iopub.execute_input":"2023-11-07T00:45:07.271958Z","iopub.status.idle":"2023-11-07T00:45:07.29397Z","shell.execute_reply.started":"2023-11-07T00:45:07.271921Z","shell.execute_reply":"2023-11-07T00:45:07.292882Z"},"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-11-07T00:45:07.295691Z","iopub.execute_input":"2023-11-07T00:45:07.296159Z","iopub.status.idle":"2023-11-07T00:45:09.974173Z","shell.execute_reply.started":"2023-11-07T00:45:07.2961Z","shell.execute_reply":"2023-11-07T00:45:09.973347Z"},"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-11-07T00:45:09.975787Z","iopub.execute_input":"2023-11-07T00:45:09.976251Z","iopub.status.idle":"2023-11-07T00:45:10.023655Z","shell.execute_reply.started":"2023-11-07T00:45:09.976223Z","shell.execute_reply":"2023-11-07T00:45:10.022679Z"},"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","metadata":{"execution":{"iopub.status.busy":"2023-11-07T00:45:10.027969Z","iopub.execute_input":"2023-11-07T00:45:10.028294Z","iopub.status.idle":"2023-11-07T00:45:11.322458Z","shell.execute_reply.started":"2023-11-07T00:45:10.028269Z","shell.execute_reply":"2023-11-07T00:45:11.321599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install fairseq\nfrom typing import Tuple, Dict, Optional\nfrom fairseq.models import FairseqEncoder,FairseqIncrementalDecoder\nfrom fairseq.modules import TransformerDecoderLayer,TransformerEncoderLayer\nimport torch\n\n\nclass Namespace(object):\n    def __init__(self, adict):\n        self.__dict__.update(adict)\n\n\nclass TransformerDecode(FairseqIncrementalDecoder):\n    def __init__(self, dim, ff_dim, num_head, num_layer):\n        super().__init__({})\n\n        self.layer = nn.ModuleList([\n            TransformerDecoderLayer(Namespace({\n                'decoder_embed_dim': dim,\n                'decoder_attention_heads': num_head,\n                'attention_dropout': 0.1,\n                'dropout': 0.1,\n                'decoder_normalize_before': True,\n                'decoder_ffn_embed_dim': ff_dim,\n                # 'decoder_learned_pos': True,\n                # 'cross_self_attention': True,\n                # 'activation-fn': 'gelu',\n            })) for i in range(num_layer)\n        ])\n        self.layer_norm = nn.LayerNorm(dim)\n\n    def forward(self, x, mem, x_mask):\n        # print('my TransformerDecode forward()')\n        for layer in self.layer:\n            x = layer(x, mem, self_attn_mask=x_mask)[0]\n        x = self.layer_norm(x)\n        return x  # T x B x C\n\n    # def forward_one(self, x, mem, incremental_state):\n    def forward_one(self,\n                    x: torch.Tensor,\n                    mem: torch.Tensor,\n                    incremental_state: Optional[Dict[str, Dict[str, Optional[torch.Tensor]]]]\n                    ) -> torch.Tensor:\n        x = x[-1:]\n        for layer in self.layer:\n            x = layer(x, mem, incremental_state=incremental_state)[0]\n        x = self.layer_norm(x)\n        return x\n","metadata":{"execution":{"iopub.status.busy":"2023-11-07T00:45:11.323518Z","iopub.execute_input":"2023-11-07T00:45:11.3238Z","iopub.status.idle":"2023-11-07T00:46:52.411699Z","shell.execute_reply.started":"2023-11-07T00:45:11.323776Z","shell.execute_reply":"2023-11-07T00:46:52.410657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import Tensor\n\nmodel_name = 'swin_small_patch4_window7_224'\n\nclass PositionEncode1D(torch.nn.Module):\n    def __init__(self, dim, max_length):\n        super().__init__()\n        assert (dim % 2 == 0)\n        self.max_length = max_length\n\n        d = torch.exp(torch.arange(0., dim, 2) * (-math.log(10000.0) / dim))\n        position = torch.arange(0., max_length).unsqueeze(1)\n        pos = torch.zeros(1, max_length, dim)\n        pos[0, :, 0::2] = torch.sin(position * d)\n        pos[0, :, 1::2] = torch.cos(position * d)\n        self.register_buffer('pos', pos)\n\n    def forward(self, x):\n        batch_size, T, dim = x.shape\n        x = x + self.pos[:, :T]\n        return x\n\n\nclass Encoder(nn.Module):\n    \"\"\"编码器类\n    这里是使用 tit 作为编码器\n    \"\"\"\n\n    def __init__(self, encoder_dim):\n        super(Encoder, self).__init__()\n        # 别人用的这个\n        # self.e = swin_base_patch4_window12_384_in22k(pretrained=True)\n\n        # 生成网络加载预训练权重\n        self.swin_t = timm.create_model(\n            model_name,\n            pretrained=True,\n            # drop_path_rate=0.2\n        )\n        self.swin_t.head = nn.Identity()\n        # 因为这里的swin_t初始的conv的out_channels=128, 所以计算到最后c为1024\n        print(self.swin_t)\n        self.fc = nn.Linear(768, encoder_dim)\n\n    def forward(self, x):\n        \"\"\"\n        :param x: shape=[b,7,7,1024] \n        \"\"\"\n        x = self.swin_t(x)\n\n        b, h, w, c = x.shape\n        x = x.reshape(b, -1, c)\n        return self.fc(x)\n\n\nclass Decoder(nn.Module):\n    def __init__(self, vocab_size, embed_dim, max_length,\n                 num_head=8, ff_dim=1024, num_layer=3):\n        super().__init__()\n\n        self.embed_dim = embed_dim\n\n        # 解码器\n        self.embed = torch.nn.Embedding(vocab_size, embed_dim)\n        self.pos = PositionEncode1D(embed_dim, max_length)\n\n        self.text_decode = TransformerDecode(embed_dim, ff_dim, num_head, num_layer)\n        self.fc = nn.Linear(embed_dim, vocab_size)\n\n    def forward(self, encoder_out, seq):\n        \"\"\"训练阶段：前向传播\n        \"\"\"\n        device = seq.device\n        b = seq.size(0)\n        encoder_out = encoder_out.permute(1, 0, 2).contiguous()\n\n        text_embed = self.embed(seq)\n        text_embed = self.pos(text_embed).permute(1, 0, 2).contiguous()\n\n        text_mask = np.triu(np.ones((seq.size(-1), seq.size(-1))), k=1).astype(np.uint8)\n        text_mask = torch.autograd.Variable(torch.from_numpy(text_mask) == 1).to(device)\n\n        #\n        x = self.text_decode(text_embed, encoder_out, text_mask)\n        x = x.permute(1, 0, 2).contiguous()\n\n        return self.fc(x)\n\n    def predict(self, encoder_out, max_len, start_token, end_token, pad_token):\n        \"\"\"预测阶段：前向传播\n        \"\"\"\n        device = encoder_out.device\n        batch_size = len(encoder_out)\n\n        # 同上\n        image_embed = encoder_out.permute(1, 0, 2).contiguous()\n\n        # b*n 填充 <pad>\n        token = torch.full((batch_size, max_len), pad_token, dtype=torch.long).to(device)\n        # 获取输出向量的位置向量\n        text_pos = self.pos.pos\n        # 第一个设置为 <sos>\n        token[:, 0] = start_token\n\n        # fast version\n        # https://github.com/pytorch/fairseq/blob/21b8fb5cb1a773d0fdc09a28203fe328c4d2b94b/fairseq/sequence_generator.py#L245-L247\n        if 1:\n            incremental_state = torch.jit.annotate(\n                Dict[str, Dict[str, Optional[Tensor]]],\n                torch.jit.annotate(Dict[str, Dict[str, Optional[Tensor]]], {}),\n            )\n            for t in range(max_len - 1):\n                # 目前 token 的最后一个值\n                last_token = token[:, t]\n                # 向量嵌入\n                text_embed = self.embed(last_token)\n                # 加上位置向量\n                text_embed = text_embed + text_pos[:, t]  #\n                # b*text_dim -> 1*b*text_dim\n                text_embed = text_embed.reshape(1, batch_size, self.embed_dim)\n                # 得到下一个向量 1*b*text_dim\n                x = self.text_decode.forward_one(text_embed, image_embed, incremental_state)\n                # b*text_dim\n                x = x.reshape(batch_size, self.embed_dim)\n                # b*vocab_size\n                l = self.fc(x)\n                # 以最大的作为预测\n                k = torch.argmax(l, -1)\n                token[:, t + 1] = k\n\n                # 遇到 <eos> 和 <pad> 停止预测\n                if ((k == end_token) | (k == pad_token)).all():\n                    break\n\n        # 返回除了 <sos> 之外的序列\n        predict = token[:, 1:]\n        return predict\n\n\nclass InCHImgAnalyzer(nn.Module):\n    def __init__(self, encoder_dim, vocab_size, embed_dim, max_length,\n                 num_head=8, ff_dim=1024, num_layer=3):\n        super().__init__()\n        self.encoder = Encoder(encoder_dim)\n        self.decoder = Decoder(\n            vocab_size, embed_dim, max_length,\n            num_head, ff_dim, num_layer\n        )\n\n    def forward(self, img, seq):\n        encoder_out = self.encoder(img)\n        return self.decoder(encoder_out, seq)\n\n    def predict(self, img, max_len, start_token, end_token, pad_token):\n        encoder_out = self.encoder(img)\n        return self.decoder.predict(encoder_out, max_len, start_token, end_token, pad_token)\n","metadata":{"execution":{"iopub.status.busy":"2023-11-07T00:46:52.413169Z","iopub.execute_input":"2023-11-07T00:46:52.413816Z","iopub.status.idle":"2023-11-07T00:46:52.44194Z","shell.execute_reply.started":"2023-11-07T00:46:52.41378Z","shell.execute_reply":"2023-11-07T00:46:52.440753Z"},"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-11-07T00:46:52.443935Z","iopub.execute_input":"2023-11-07T00:46:52.444287Z","iopub.status.idle":"2023-11-07T00:46:52.470218Z","shell.execute_reply.started":"2023-11-07T00:46:52.444261Z","shell.execute_reply":"2023-11-07T00:46:52.469363Z"},"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-11-07T00:46:52.471419Z","iopub.execute_input":"2023-11-07T00:46:52.471699Z","iopub.status.idle":"2023-11-07T00:46:52.491281Z","shell.execute_reply.started":"2023-11-07T00:46:52.471674Z","shell.execute_reply":"2023-11-07T00:46:52.490294Z"},"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-11-07T00:46:52.492441Z","iopub.execute_input":"2023-11-07T00:46:52.492721Z","iopub.status.idle":"2023-11-07T00:46:52.507701Z","shell.execute_reply.started":"2023-11-07T00:46:52.492697Z","shell.execute_reply":"2023-11-07T00:46:52.506649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(TrainConfig.DEVICE)","metadata":{"execution":{"iopub.status.busy":"2023-11-07T00:46:52.508818Z","iopub.execute_input":"2023-11-07T00:46:52.509125Z","iopub.status.idle":"2023-11-07T00:46:52.52509Z","shell.execute_reply.started":"2023-11-07T00:46:52.509098Z","shell.execute_reply":"2023-11-07T00:46:52.523908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# main_preprocess()","metadata":{"execution":{"iopub.status.busy":"2023-11-07T00:46:52.526257Z","iopub.execute_input":"2023-11-07T00:46:52.5266Z","iopub.status.idle":"2023-11-07T00:46:52.53582Z","shell.execute_reply.started":"2023-11-07T00:46:52.52655Z","shell.execute_reply":"2023-11-07T00:46:52.535034Z"},"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-11-07T00:46:52.536936Z","iopub.execute_input":"2023-11-07T00:46:52.537259Z","iopub.status.idle":"2023-11-07T00:46:52.566163Z","shell.execute_reply.started":"2023-11-07T00:46:52.537234Z","shell.execute_reply":"2023-11-07T00:46:52.565248Z"},"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-11-07T00:46:52.567668Z","iopub.execute_input":"2023-11-07T00:46:52.568395Z","iopub.status.idle":"2023-11-07T00:46:52.582688Z","shell.execute_reply.started":"2023-11-07T00:46:52.568361Z","shell.execute_reply":"2023-11-07T00:46:52.58175Z"},"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)\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        label_text_pred = tokenizer.predict_captions(preds.detach().cpu().numpy())\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        save_net(net, optimizer, scheduler, Config.PATH.LAST_WEIGHT_PATH)\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    LOG.info(f'Net Saved, path: {filepath}')\n","metadata":{"execution":{"iopub.status.busy":"2023-11-07T00:46:52.584314Z","iopub.execute_input":"2023-11-07T00:46:52.584711Z","iopub.status.idle":"2023-11-07T00:46:52.613754Z","shell.execute_reply.started":"2023-11-07T00:46:52.584676Z","shell.execute_reply":"2023-11-07T00:46:52.612734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main_train():\n    seed_torch(Config.TRAIN.SEED)\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 = InCHImgAnalyzer(\n        encoder_dim=Config.TRAIN.ENCODER_DIM,\n        vocab_size=len(tokenizer),\n        embed_dim=Config.TRAIN.EMBED_DIM,\n        max_length=Config.TRAIN.MAX_LEN,\n        num_head=Config.TRAIN.N_HEAD,\n        ff_dim=Config.TRAIN.FF_DIM,\n        num_layer=Config.TRAIN.NUM_LAYER\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-11-07T00:46:52.615083Z","iopub.execute_input":"2023-11-07T00:46:52.615647Z","iopub.status.idle":"2023-11-07T00:46:52.629356Z","shell.execute_reply.started":"2023-11-07T00:46:52.615614Z","shell.execute_reply":"2023-11-07T00:46:52.628495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"main_train()","metadata":{"execution":{"iopub.status.busy":"2023-11-07T00:46:52.63049Z","iopub.execute_input":"2023-11-07T00:46:52.630767Z","iopub.status.idle":"2023-11-07T00:48:03.119013Z","shell.execute_reply.started":"2023-11-07T00:46:52.63074Z","shell.execute_reply":"2023-11-07T00:48:03.117886Z"},"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-11-07T00:48:03.180444Z","iopub.execute_input":"2023-11-07T00:48:03.180925Z","iopub.status.idle":"2023-11-07T00:48:03.261173Z","shell.execute_reply.started":"2023-11-07T00:48:03.180889Z","shell.execute_reply":"2023-11-07T00:48:03.260127Z"},"trusted":true},"execution_count":null,"outputs":[]}]}