{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":7330212,"sourceType":"datasetVersion","datasetId":4245149,"isSourceIdPinned":false}],"dockerImageVersionId":30626,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install pytorch_toolbelt\n!pip install nystrom-attention","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-03T11:08:48.685762Z","iopub.execute_input":"2024-01-03T11:08:48.686065Z","iopub.status.idle":"2024-01-03T11:09:14.958598Z","shell.execute_reply.started":"2024-01-03T11:08:48.68604Z","shell.execute_reply":"2024-01-03T11:09:14.957529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!jupyter notebook --NotebookApp.iopub_data_rate_limit=1.0e14","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-03T11:09:14.960746Z","iopub.execute_input":"2024-01-03T11:09:14.961062Z","iopub.status.idle":"2024-01-03T11:09:25.694297Z","shell.execute_reply.started":"2024-01-03T11:09:14.961033Z","shell.execute_reply":"2024-01-03T11:09:25.693032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#----> path\nfrom pathlib import Path\nimport glob\nimport os\n\n#----> utils\nimport sys\nimport numpy as np\nimport inspect # 查看python 类的参数和模块、函数代码\nimport importlib # In order to dynamically import the library\nimport random\nimport pandas as pd\nfrom sklearn.model_selection import StratifiedShuffleSplit, StratifiedGroupKFold\nimport matplotlib.pyplot as plt\nimport warnings\n%matplotlib inline\nwarnings.filterwarnings('ignore')\n\n#----> pytorch\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchmetrics\nfrom einops import rearrange\nfrom nystrom_attention import NystromAttention\nfrom torch.utils.data import random_split, DataLoader, Dataset\nfrom torchvision.datasets import MNIST\nfrom torchvision import transforms\n\n# Loss\nfrom pytorch_toolbelt import losses as L\n\n# Optim\nimport math\nimport torch.optim as optim\nfrom torch.optim.optimizer import Optimizer\nfrom collections import defaultdict\ntry:\n    from apex.optimizers import FusedNovoGrad, FusedAdam, FusedLAMB, FusedSGD\n    has_apex = True\nexcept ImportError:\n    has_apex = False\n\n\n#----> pytorch_lightning\nimport pytorch_lightning as pl\nfrom pytorch_lightning import Trainer\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nfrom pytorch_lightning.callbacks.early_stopping import EarlyStopping\nfrom pytorch_lightning import loggers as pl_loggers","metadata":{"execution":{"iopub.status.busy":"2024-01-03T11:09:25.695872Z","iopub.execute_input":"2024-01-03T11:09:25.696248Z","iopub.status.idle":"2024-01-03T11:09:33.082017Z","shell.execute_reply.started":"2024-01-03T11:09:25.696214Z","shell.execute_reply":"2024-01-03T11:09:33.081197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONFIG = {\n    # Config\n    \"config\": 'UBC',\n    \n    # Generate\n    \"seed\": 42,\n    \"devices\": 1,\n    \"device\": \"cuda\" if torch.cuda.is_available() else \"cpu\",\n    \"benchmark\": True,\n    \"precision\": 32,\n    \"epochs\": 100,\n    \"grad_acc\": 2,\n    \"patience\": 5,\n    \"server\": \"train\", #train #test\n    \"log_path\": \"logs/\",\n    \"save_top_k\": 1,\n    \"test_size\": 0.2,\n    \n    # Data\n    \"dataset_name\": \"UBC_data\",\n    \"data_shuffle\": True,\n    \"data_dir\": \"/kaggle/input/feature-scale0-25/pt_files\",\n    \"dataset_dir\": \"/kaggle/working/dataset_csv/UBC16\",\n    \"fold\": 0,\n    \"nfold\": 4,\n    \n    \"train_dataloader\": {\n           \"batch_size\": 1, \n            \"num_workers\": 8,\n       },\n\n    \"test_dataloader\": {\n            \"batch_size\": 1,\n            \"num_workers\": 8,\n        },\n    \n    # Model\n    \"name\": \"MMIL\",\n    \"n_classes\": 5,\n    \"mask_ratio\": 0.3,   # mask的概率\n    \"D_feat\": 1024,    # 输入特征维度\n    \"D_inner\": 512,   # 中间特征维度\n    \"mode\": \"random\",\n    \"num_subbags\": 10,  # 调参 子包数量\n    \"ape\": True,\n    \"num_layers\": 2,  # 调参\n\n    # Optimizer\n    \"opt\": \"lookahead_radam\",\n    \"lr\": 1e-4,\n    \"opt_eps\": None, \n    \"opt_betas\": None,\n    \"momentum\": None, \n    \"weight_decay\": 1e-5,\n    \n    # Loss\n    \"base_loss\": \"focal\",\n}\n\nlabel_encoder = {\n    \"MC\": 0,\n    \"EC\": 1,\n    \"CC\": 2,\n    \"HGSC\": 3,\n    \"LGSC\": 4,\n}\n\nlabel_decoder = {v:k for k,v in label_encoder.items()}\n\ndef transform_data(s):\n    s['image_id'] = str(s['image_id'])\n    s['label'] = label_encoder[s['label']]\n    return s","metadata":{"execution":{"iopub.status.busy":"2024-01-03T11:13:23.456176Z","iopub.execute_input":"2024-01-03T11:13:23.456529Z","iopub.status.idle":"2024-01-03T11:13:23.466427Z","shell.execute_reply.started":"2024-01-03T11:13:23.4565Z","shell.execute_reply":"2024-01-03T11:13:23.465435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def drop_useless_img(df_data):\n    useless_list = [281, 3222, 5264, 9154, 12244, 26124, 31793, 32192, 33839, 41099, 52308, 54506, 63836, 1289, 32035]\n    df_data = df_data[ ~df_data['image_id'].isin(useless_list) ]\n    if 15583 in df_data['image_id']:\n        df_data[ df_data['image_id'] == 15583 ]['label'] = 'MC'\n    return df_data\n\n\ndef generate_dataset_csv(use_tma = True):        \n    \n    df_data = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\n    df_data = drop_useless_img(df_data)\n    df_data = df_data.apply(lambda s: transform_data(s), axis=1)\n    \n    slide_ids = os.listdir(CONFIG['data_dir'])\n    slide_ids = [os.path.splitext(slide_id)[0] for slide_id in slide_ids]\n    cnt = 0\n    for slide_id in slide_ids:\n        if slide_id in df_data['image_id'].values:\n            continue\n        else:\n            cnt += 1\n            print(f\"slide id: {slide_id} is remove.\")\n    \n    df_data = df_data[ df_data['image_id'].isin(slide_ids) ]\n    print(f\"df_data len is {len(df_data)}, slide_ids len is {len(slide_ids)}, drop ids num is {cnt}.\")\n    \n    df_tma = df_data[ df_data['is_tma'] == True ]\n    df_data = df_data[ ~df_data['image_id'].isin(df_tma['image_id'].values) ].loc[:, ['image_id', 'label']].reset_index(drop=True)\n    \n    print(f\"df_tma len is {len(df_tma)}, df_data len is {len(df_data)}.\")\n    \n    # 分层交叉验证\n    skf = StratifiedShuffleSplit(n_splits=CONFIG['nfold'], test_size=CONFIG['test_size'])\n    for fold, (train_idxs, val_idxs) in enumerate(skf.split(df_data['image_id'], df_data['label'])):\n        print(\"分层随机划分：%s %s\" % (train_idxs.shape, val_idxs.shape))\n        # 训练集\n        df_train = df_data.iloc[train_idxs, :].reset_index(drop=True)\n        # 验证集\n        df_val = df_data.iloc[val_idxs, :].reset_index(drop=True)\n        # 测试集\n        df_test = df_tma.loc[:, ['image_id', 'label']].sample(frac=1.0).reset_index(drop=True)\n        # 重新命名列名\n        columns_name = [\"train_id\", \"train_label\", \"val_id\", \"val_label\", \"test_id\", \"test_label\"]\n        df_dataset = pd.concat([df_train, df_val, df_test], axis=1, ignore_index=True).set_axis(columns_name, axis=1)\n        # 创建 dataset_csv 的存储路径\n        Path(CONFIG['dataset_dir']).mkdir(exist_ok=True, parents=True)\n        df_dataset.to_csv(os.path.join(CONFIG['dataset_dir'], f\"fold{fold}.csv\"), index=False)\n        display(df_dataset)\n    \ngenerate_dataset_csv(use_tma = True)","metadata":{"execution":{"iopub.status.busy":"2024-01-03T11:14:30.203379Z","iopub.execute_input":"2024-01-03T11:14:30.204087Z","iopub.status.idle":"2024-01-03T11:14:30.360867Z","shell.execute_reply.started":"2024-01-03T11:14:30.204056Z","shell.execute_reply":"2024-01-03T11:14:30.359824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RAdam(Optimizer):\n\n    def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=0):\n        defaults = dict(lr=lr, betas=betas, eps=eps, weight_decay=weight_decay)\n        self.buffer = [[None, None, None] for ind in range(10)]\n        super(RAdam, self).__init__(params, defaults)\n\n    def __setstate__(self, state):\n        super(RAdam, self).__setstate__(state)\n\n    def step(self, closure=None):\n\n        loss = None\n        if closure is not None:\n            loss = closure()\n\n        for group in self.param_groups:\n\n            for p in group['params']:\n                if p.grad is None:\n                    continue\n                grad = p.grad.data.float()\n                if grad.is_sparse:\n                    raise RuntimeError('RAdam does not support sparse gradients')\n\n                p_data_fp32 = p.data.float()\n\n                state = self.state[p]\n\n                if len(state) == 0:\n                    state['step'] = 0\n                    state['exp_avg'] = torch.zeros_like(p_data_fp32)\n                    state['exp_avg_sq'] = torch.zeros_like(p_data_fp32)\n                else:\n                    state['exp_avg'] = state['exp_avg'].type_as(p_data_fp32)\n                    state['exp_avg_sq'] = state['exp_avg_sq'].type_as(p_data_fp32)\n\n                exp_avg, exp_avg_sq = state['exp_avg'], state['exp_avg_sq']\n                beta1, beta2 = group['betas']\n\n                exp_avg_sq.mul_(beta2).addcmul_(1 - beta2, grad, grad)\n                exp_avg.mul_(beta1).add_(1 - beta1, grad)\n\n                state['step'] += 1\n                buffered = self.buffer[int(state['step'] % 10)]\n                if state['step'] == buffered[0]:\n                    N_sma, step_size = buffered[1], buffered[2]\n                else:\n                    buffered[0] = state['step']\n                    beta2_t = beta2 ** state['step']\n                    N_sma_max = 2 / (1 - beta2) - 1\n                    N_sma = N_sma_max - 2 * state['step'] * beta2_t / (1 - beta2_t)\n                    buffered[1] = N_sma\n\n                    # more conservative since it's an approximated value\n                    if N_sma >= 5:\n                        step_size = group['lr'] * math.sqrt(\n                            (1 - beta2_t) * (N_sma - 4) / (N_sma_max - 4) * (N_sma - 2) / N_sma * N_sma_max / (\n                                        N_sma_max - 2)) / (1 - beta1 ** state['step'])\n                    else:\n                        step_size = group['lr'] / (1 - beta1 ** state['step'])\n                    buffered[2] = step_size\n\n                if group['weight_decay'] != 0 and group['weight_decay'] is not None:\n                    p_data_fp32.add_(-group['weight_decay'] * group['lr'], p_data_fp32)\n\n                # more conservative since it's an approximated value\n                if N_sma >= 5:\n                    denom = exp_avg_sq.sqrt().add_(group['eps'])\n                    p_data_fp32.addcdiv_(-step_size, exp_avg, denom)\n                else:\n                    p_data_fp32.add_(-step_size, exp_avg)\n\n                p.data.copy_(p_data_fp32)\n\n        return loss","metadata":{"execution":{"iopub.status.busy":"2024-01-03T11:14:30.847282Z","iopub.execute_input":"2024-01-03T11:14:30.848118Z","iopub.status.idle":"2024-01-03T11:14:30.865988Z","shell.execute_reply.started":"2024-01-03T11:14:30.848082Z","shell.execute_reply":"2024-01-03T11:14:30.865013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_loggers():\n\n    log_path = CONFIG['log_path']\n    Path(log_path).mkdir(exist_ok=True, parents=True)\n    \n    log_name = CONFIG['config'] \n    version_name = CONFIG['name']\n    CONFIG['fold_log_path'] = Path(log_path) / log_name / version_name / f\"fold{CONFIG['fold']}\"  # CONFIG['fold'] 在交叉验证过程中不断改变\n    print(f\"---->Log dir: {CONFIG['fold_log_path']}\")\n    \n    #---->TensorBoard\n    tb_logger = pl_loggers.TensorBoardLogger(log_path+str(log_name),\n                                             name = version_name, version = f\"fold{CONFIG['fold']}\",\n                                             log_graph = True, default_hp_metric = False)\n    #---->CSV\n    csv_logger = pl_loggers.CSVLogger(log_path+str(log_name),\n                                      name = version_name, version = f\"fold{CONFIG['fold']}\", )\n    \n    return [tb_logger, csv_logger]\n\n\ndef load_callbacks():\n\n    Mycallbacks = []\n    # Make output path\n    output_path = CONFIG['fold_log_path']\n    output_path.mkdir(exist_ok=True, parents=True)\n\n    early_stop_callback = EarlyStopping(\n        monitor='val_loss',\n        min_delta=0.00,\n        patience=CONFIG['patience'],\n        verbose=True,\n        mode='min'\n    )\n    Mycallbacks.append(early_stop_callback)\n\n    if CONFIG['server'] == 'train':\n        Mycallbacks.append(ModelCheckpoint(monitor = 'val_loss',\n                                         dirpath = str(CONFIG['fold_log_path']),\n                                         filename = '{epoch:02d}-{val_loss:.4f}',\n                                         verbose = True,\n                                         save_last = True,\n                                         save_top_k = CONFIG['save_top_k'],\n                                         mode = 'min',\n                                         save_weights_only = True))\n    return Mycallbacks\n\n\ndef cross_entropy_torch(x, y):\n    x_softmax = [F.softmax(x[i]) for i in range(len(x))]\n    x_log = torch.tensor([torch.log(x_softmax[i][y[i]]) for i in range(len(y))])\n    loss = - torch.sum(x_log) / len(y)\n    return loss\n\n\ndef create_loss(select_loss, w1=1.0, w2=0.5):\n    conf_loss = select_loss\n    ### MulticlassJaccardLoss(classes=np.arange(11)\n    # mode = args.base_loss #BINARY_MODE \\MULTICLASS_MODE \\MULTILABEL_MODE \n    loss = None\n    if hasattr(nn, conf_loss): \n        loss = getattr(nn, conf_loss)() \n    #binary loss\n    elif conf_loss == \"focal\":\n        loss = L.CrossEntropyFocalLoss()\n    elif conf_loss == \"jaccard\":\n        loss = L.BinaryJaccardLoss()\n    elif conf_loss == \"jaccard_log\":\n        loss = L.BinaryJaccardLoss()\n    elif conf_loss == \"dice\":\n        loss = L.BinaryDiceLoss()\n    elif conf_loss == \"dice_log\":\n        loss = L.BinaryDiceLogLoss()\n    elif conf_loss == \"dice_log\":\n        loss = L.BinaryDiceLogLoss()\n    elif conf_loss == \"bce+lovasz\":\n        loss = L.JointLoss(BCEWithLogitsLoss(), L.BinaryLovaszLoss(), w1, w2)\n    elif conf_loss == \"lovasz\":\n        loss = L.BinaryLovaszLoss()\n    elif conf_loss == \"bce+jaccard\":\n        loss = L.JointLoss(BCEWithLogitsLoss(), L.BinaryJaccardLoss(), w1, w2)\n    elif conf_loss == \"bce+log_jaccard\":\n        loss = L.JointLoss(BCEWithLogitsLoss(), L.BinaryJaccardLogLoss(), w1, w2)\n    elif conf_loss == \"bce+log_dice\":\n        loss = L.JointLoss(BCEWithLogitsLoss(), L.BinaryDiceLogLoss(), w1, w2)\n    elif conf_loss == \"reduced_focal\":\n        loss = L.BinaryFocalLoss(reduced=True)\n    else:\n        assert False and \"Invalid loss\"\n        raise ValueError\n    return loss","metadata":{"execution":{"iopub.status.busy":"2024-01-03T11:14:31.331169Z","iopub.execute_input":"2024-01-03T11:14:31.33154Z","iopub.status.idle":"2024-01-03T11:14:31.348688Z","shell.execute_reply.started":"2024-01-03T11:14:31.331509Z","shell.execute_reply":"2024-01-03T11:14:31.347841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UBCData(Dataset):\n    def __init__(self, state=None):\n        # Set all input args as attributes\n        self.__dict__.update(locals())\n\n        #---->data and label\n        self.nfolds = CONFIG['nfold']\n        self.fold = CONFIG['fold']\n        self.feature_dir = CONFIG['data_dir']\n        self.csv_dir = os.path.join(CONFIG['dataset_dir'], f'fold{self.fold}.csv')\n        self.slide_data = pd.read_csv(self.csv_dir)\n\n        #---->order\n        self.shuffle = CONFIG['data_shuffle']\n\n        #---->split dataset\n        if state == 'train':\n            self.data = self.slide_data.loc[:, 'train_id'].dropna()\n            self.label = self.slide_data.loc[:, 'train_label'].dropna()\n        if state == 'val':\n            self.data = self.slide_data.loc[:, 'val_id'].dropna()\n            self.label = self.slide_data.loc[:, 'val_label'].dropna()\n        if state == 'test':\n            self.data = self.slide_data.loc[:, 'test_id'].dropna()\n            self.label = self.slide_data.loc[:, 'test_label'].dropna()\n\n\n    def __len__(self):\n        return len(self.data)\n\n    \n    def __getitem__(self, idx):\n        slide_id = int(self.data[idx])\n        label = int(self.label[idx]) # 数据中存在 NAN 会自动转化为浮点型，因此转化成整数\n        pt_path = os.path.join(CONFIG['data_dir'], str(slide_id) + '.pt')\n#         print(pt_path)\n        features = torch.load(pt_path)\n\n        #----> shuffle\n        if self.shuffle == True:\n            index = [x for x in range(features.shape[0])]\n            random.shuffle(index)\n            features = features[index]\n\n\n        return features, label\n\n    \nmy_dataset = UBCData(state='test')\nfor i in range(5):\n    feature, label = my_dataset[i]\n    print(feature.shape, label)","metadata":{"execution":{"iopub.status.busy":"2024-01-03T11:14:31.654236Z","iopub.execute_input":"2024-01-03T11:14:31.65509Z","iopub.status.idle":"2024-01-03T11:14:31.901776Z","shell.execute_reply.started":"2024-01-03T11:14:31.655056Z","shell.execute_reply":"2024-01-03T11:14:31.900471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataInterface(pl.LightningDataModule):\n\n    def __init__(self, train_batch_size=64, train_num_workers=8, test_batch_size=1, test_num_workers=1,dataset_name=None, **kwargs):\n        \"\"\"[summary]\n\n        Args:\n            batch_size (int, optional): [description]. Defaults to 64.\n            num_workers (int, optional): [description]. Defaults to 8.\n            dataset_name (str, optional): [description]. Defaults to ''.\n        \"\"\"        \n        super().__init__()\n\n        self.train_batch_size = train_batch_size\n        self.train_num_workers = train_num_workers\n        self.test_batch_size = test_batch_size\n        self.test_num_workers = test_num_workers\n        self.dataset_name = dataset_name\n        self.kwargs = kwargs\n        self.load_data_module()\n\n \n    def prepare_data(self):\n        pass\n\n    \n    def setup(self, stage=None):\n        # 2. how to split, argument\n        \"\"\"  \n        - count number of classes\n\n        - build vocabulary\n\n        - perform train/val/test splits\n\n        - apply transforms (defined explicitly in your datamodule or assigned in init)\n        \"\"\"\n        # Assign train/val datasets for use in dataloaders\n        if stage == 'fit' or stage is None:\n            self.train_dataset = self.instancialize(state='train')\n            self.val_dataset = self.instancialize(state='val')\n \n\n        # Assign test dataset for use in dataloader(s)\n        if stage == 'test' or stage is None:\n            self.test_dataset = self.instancialize(state='test')\n\n\n    def train_dataloader(self):\n        return DataLoader(self.train_dataset, batch_size=self.train_batch_size, num_workers=self.train_num_workers, shuffle=True)\n\n    \n    def val_dataloader(self):\n        return DataLoader(self.val_dataset, batch_size=self.train_batch_size, num_workers=self.train_num_workers, shuffle=False)\n\n    \n    def test_dataloader(self):\n        return DataLoader(self.test_dataset, batch_size=self.test_batch_size, num_workers=self.test_num_workers, shuffle=False)\n    \n    \n    def load_data_module(self):\n        camel_name =  ''.join([i.capitalize() for i in (self.dataset_name).split('_')])\n        try:\n            self.data_module = UBCData\n        except:\n            raise ValueError(\n                'Invalid Dataset File Name or Invalid Class Name!')\n    \n    \n    def instancialize(self, **other_args):\n        \"\"\" Instancialize a model using the corresponding parameters\n            from self.hparams dictionary. You can also input any args\n            to overwrite the corresponding value in self.kwargs.\n        \"\"\"\n        class_args = inspect.getargspec(self.data_module.__init__).args[1:]\n        inkeys = self.kwargs.keys()\n        args1 = {}\n        for arg in class_args:\n            if arg in inkeys:\n                args1[arg] = self.kwargs[arg]\n        args1.update(other_args)\n        return self.data_module(**args1)\n\n    \nDataInterface_dict = {\n                        'train_batch_size': CONFIG['train_dataloader']['batch_size'],\n                        'train_num_workers': CONFIG['train_dataloader']['num_workers'],\n                        'test_batch_size': CONFIG['test_dataloader']['batch_size'],\n                        'test_num_workers': CONFIG['test_dataloader']['num_workers'],\n                        'dataset_name': CONFIG['dataset_name'],\n                    }\nmy_dm = DataInterface(**DataInterface_dict)\nmy_dm.setup(stage='fit')\nprint(len(my_dm.val_dataloader()))\nfor feature, label in my_dm.val_dataloader():\n    print(feature.shape, label)\n    break","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-03T11:14:31.994979Z","iopub.execute_input":"2024-01-03T11:14:31.995721Z","iopub.status.idle":"2024-01-03T11:14:32.284867Z","shell.execute_reply.started":"2024-01-03T11:14:31.995676Z","shell.execute_reply":"2024-01-03T11:14:32.283693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def cat_msg2cluster_group(x_groups,msg_tokens):\n    x_groups_cated = []\n    for x in x_groups:\n        x = x.unsqueeze(dim=0)\n        try:\n            temp = torch.cat((msg_tokens,x),dim=2)\n        except Exception as e:\n            print('Error when cat msg tokens to sub-bags')\n        x_groups_cated.append(temp)\n\n    return x_groups_cated\n\n\n\ndef split_array(array, m):\n    n = len(array)\n    indices = np.random.choice(n, n, replace=False)\n    split_indices = np.array_split(indices, m)  \n\n    result = []\n    for indices in split_indices:\n        result.append(array[indices])\n\n    return result\n\n\n\nclass grouping:\n\n    def __init__(self,groups_num,max_size=1e10):\n        self.groups_num = groups_num\n        self.max_size = int(max_size) # Max lenth 4300 for 24G RTX3090\n        \n    \n    def indicer(self, labels):\n        indices = []\n        groups_num = len(set(labels))\n        for i in range(groups_num):\n            temp = np.argwhere(labels==i).squeeze()\n            indices.append(temp)\n        return indices\n    \n    def make_subbags(self, idx, features):\n        index = idx\n        features_group = []\n        for i in range(len(index)):\n            member_size = (index[i].size)\n            if member_size > self.max_size:\n                index[i] = np.random.choice(index[i],size=self.max_size,replace=False)\n            temp = features[index[i]]\n            temp = temp.unsqueeze(dim=0)\n            features_group.append(temp)\n            \n        return features_group\n        \n    \n    def embedding_grouping(self,features):\n        features = features.squeeze()\n        k = KMeans(n_clusters=self.groups_num, random_state=0,n_init='auto').fit(features.cpu().detach().numpy())\n        indices = self.indicer(k.labels_)\n        features_group = self.make_subbags(indices,features)\n\n        return features_group\n    \n    def random_grouping(self, features):\n        B, N, C = features.shape\n        features = features.squeeze()\n        indices = split_array(np.array(range(int(N))),self.groups_num)\n        features_group = self.make_subbags(indices,features)\n        \n        return features_group\n\n    \n\nclass Attention(nn.Module):\n    def __init__(self, dim, heads = 8, dim_head = 64, dropout = 0.):\n        super().__init__()\n        inner_dim = dim_head *  heads\n        project_out = not (heads == 1 and dim_head == dim)\n\n        self.heads = heads\n        self.scale = dim_head ** -0.5\n\n        self.attend = nn.Softmax(dim = -1)\n        self.dropout = nn.Dropout(dropout)\n\n        self.to_qkv = nn.Linear(dim, inner_dim * 3, bias = False)\n\n        self.to_out = nn.Sequential(\n            nn.Linear(inner_dim, dim),\n            nn.Dropout(dropout)\n        ) if project_out else nn.Identity()\n\n        \n    def forward(self, x):\n        #x = x.squeeze(dim=0)\n        qkv = self.to_qkv(x).chunk(3, dim = -1)\n        q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h = self.heads), qkv)\n\n        dots = torch.matmul(q, k.transpose(-1, -2)) * self.scale\n\n        attn = self.attend(dots)\n        attn = self.dropout(attn)\n\n        out = torch.matmul(attn, v)\n        out = rearrange(out, 'b h n d -> b n (h d)')\n        return self.to_out(out)\n    \n    \nclass AttenLayer(nn.Module):\n    def __init__(self,dim,heads=8,dim_head=64,dropout=0.1,attn_mode='normal'):\n        super(AttenLayer, self).__init__()\n        self.dim = dim\n        self.heads = heads\n        self.dim_head = dim_head\n        self.dropout = dropout\n        self.mode = attn_mode\n        self.attn = Attention(self.dim,heads=self.heads,dim_head=self.dim_head,dropout=self.dropout)\n    def forward(self,x):\n        return x + self.attn(x)\n\n    \n    \nclass NyAttenLayer(nn.Module):\n    def __init__(self,dim,heads=8,dim_head=64,dropout=0.1):\n        super(NyAttenLayer, self).__init__()\n        self.dim = dim\n        self.heads = heads\n        self.dim_head = dim_head\n        self.dropout = dropout\n        self.attn = NystromAttention(\n            dim = dim,\n            dim_head = dim//8,\n            heads = 8,\n            num_landmarks = dim//2,    # number of landmarks\n            pinv_iterations = 6,    # number of moore-penrose iterations for approximating pinverse. 6 was recommended by the paper\n            residual = True,         # whether to do an extra residual with the value or not. supposedly faster convergence if turned on\n            dropout=0.1\n        )\n    def forward(self,x):\n        return x + self.attn(x)\n\n    \nclass GroupsAttenLayer(nn.Module):\n    def __init__(self,dim,heads=8,dim_head=64,dropout=0.1,attn_mode='nystrom'):\n        super(GroupsAttenLayer, self).__init__()\n        self.dim = dim\n        self.heads = heads\n        self.dim_head = dim_head\n        self.dropout = dropout\n        if attn_mode == 'nystrom':\n            self.AttenLayer = NyAttenLayer(dim =self.dim,heads=self.heads,dim_head=self.dim_head,dropout=self.dropout)\n        else:\n            self.AttenLayer = AttenLayer(dim =self.dim,heads=self.heads,dim_head=self.dim_head,dropout=self.dropout)\n\n    def forward(self,x_groups,mask_ratio=0):\n        group_after_attn = []\n        r = int(len(x_groups) * (1-mask_ratio))\n        x_groups_masked = random.sample(x_groups, k=r)\n        for x in x_groups_masked:\n            x = x.squeeze(dim=0)\n            temp = self.AttenLayer(x).unsqueeze(dim=0)\n            group_after_attn.append(temp)\n        return group_after_attn\n\n    \n\nclass GroupsMSGAttenLayer(nn.Module):\n    def __init__(self,dim,heads=8,dim_head=64,dropout=0.1):\n        super().__init__()\n        self.dim = dim\n        self.heads = heads\n        self.dim_head = dim_head\n        self.dropout = dropout\n        self.AttenLayer = AttenLayer(dim =self.dim,heads=self.heads,dim_head=self.dim_head,dropout=self.dropout)\n    def forward(self,data):\n        msg_cls, x_groups, msg_tokens_num = data\n        groups_num = len(x_groups)\n        msges = torch.zeros(size=(1,1,groups_num*msg_tokens_num,self.dim)).to(msg_cls.device)\n        for i in range(groups_num):\n            msges[:,:,i*msg_tokens_num:(i+1)*msg_tokens_num,:] = x_groups[i][:,:,0:msg_tokens_num]\n        msges = torch.cat((msg_cls,msges),dim=2).squeeze(dim=0)\n        msges = self.AttenLayer(msges).unsqueeze(dim=0)\n        msg_cls = msges[:,:,0].unsqueeze(dim=0)\n        msges = msges[:,:,1:]\n        for i in range(groups_num):\n            x_groups[i] = torch.cat((msges[:,:,i*msg_tokens_num:(i+1)*msg_tokens_num],x_groups[i][:,:,msg_tokens_num:]),dim=2)\n        data = msg_cls, x_groups, msg_tokens_num\n        return data\n    \n    \nclass BasicLayer(nn.Module):\n    def __init__(self,dim):\n        super().__init__()\n        self.GroupsAttenLayer = GroupsAttenLayer(dim=dim)\n        self.GroupsMSGAttenLayer = GroupsMSGAttenLayer(dim=dim)\n    def forward(self,data,mask_ratio):\n        msg_cls, x_groups, msg_tokens_num = data\n        x_groups = self.GroupsAttenLayer(x_groups,mask_ratio)\n        data = (msg_cls, x_groups, msg_tokens_num)\n        data = self.GroupsMSGAttenLayer(data)\n        return data\n\n    \n\nclass MultipleMILTransformer(nn.Module):\n    def __init__(self):\n        super(MultipleMILTransformer, self).__init__()\n        self.in_chan = CONFIG['D_feat']\n        self.embed_dim = CONFIG['D_inner']\n        self.n_classes = CONFIG['n_classes']\n        self.mode = CONFIG['mode']\n        self.num_subbags = CONFIG['num_subbags']\n        self.ape = CONFIG['ape']\n        self.num_layers = CONFIG['num_layers']\n        \n        self.fc1 = nn.Linear(self.in_chan, self.embed_dim)\n        self.fc2 = nn.Linear(self.embed_dim, self.n_classes)\n        self.msg_tokens_num = 1 # self.args.num_msg \n        self.msgcls_token = nn.Parameter(torch.randn(1,1,1,self.embed_dim))\n        \n        \n        #---> make sub-bags\n        print('try to group seq to ',self.num_subbags)\n        self.grouping = grouping(self.num_subbags,max_size=4300)\n        if self.mode == 'random':\n            self.grouping_features = self.grouping.random_grouping\n        elif self.mode == 'embed':\n            self.grouping_features = self.grouping.embedding_grouping\n\n        self.msg_tokens = nn.Parameter(torch.zeros(1, 1, 1, self.embed_dim))\n        \n        self.cat_msg2cluster_group = cat_msg2cluster_group\n        \n        if self.ape:\n                self.absolute_pos_embed = nn.Parameter(torch.zeros(1, 1, self.embed_dim))\n        \n        \n        #--->build layers\n        self.layers = nn.ModuleList()\n        for i_layer in range(self.num_layers):\n            layer = BasicLayer(dim=self.embed_dim)\n            self.layers.append(layer)\n\n    def head(self,x):\n        logits = self.fc2(x)\n        Y_hat = torch.argmax(logits, dim=1)\n        Y_prob = F.softmax(logits, dim=1)\n        results_dict = {'logits': logits, 'Y_prob': Y_prob, 'Y_hat': Y_hat}\n        return results_dict\n\n\n    def forward(self,x, coords=False,mask_ratio=0):\n        #---> init\n        x = self.fc1(x)\n\n        if self.ape:\n            x = x + self.absolute_pos_embed.expand(1,x.shape[1],self.embed_dim)\n        if self.mode == 'coords' or self.mode == 'idx':\n            x_groups = self.grouping_features(coords,x) \n        else:\n            x_groups = self.grouping_features(x)\n            \n        msg_tokens = self.msg_tokens.expand(1,1,self.msg_tokens_num,-1)\n        msg_cls = self.msgcls_token\n        x_groups = self.cat_msg2cluster_group(x_groups,msg_tokens)\n        data = (msg_cls, x_groups, self.msg_tokens_num)\n        #---> feature forward\n        for i in range(len(self.layers)):\n            if i == 0:\n                mr = mask_ratio\n                data = self.layers[i](data,mr)\n            else:\n                mr = 0\n                data = self.layers[i](data,mr)\n        #---> head\n        msg_cls, _, _ = data\n        msg_cls = msg_cls.view(1,self.embed_dim)\n        results_dict = self.head(msg_cls)\n        #print(results_dict)\n\n        return results_dict\n    \n    \nmilnet = MultipleMILTransformer()\nimg_feature = torch.rand(1, 100, 1024)\nresults_dict = milnet(img_feature)\nprint(results_dict['logits'].shape)\nprint(results_dict['Y_prob'].shape)\nprint(results_dict['Y_hat'].shape)","metadata":{"execution":{"iopub.status.busy":"2024-01-03T11:14:32.781715Z","iopub.execute_input":"2024-01-03T11:14:32.782068Z","iopub.status.idle":"2024-01-03T11:14:33.327454Z","shell.execute_reply.started":"2024-01-03T11:14:32.782038Z","shell.execute_reply":"2024-01-03T11:14:33.326278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sklearn.metrics\n\nclass ModelInterface(pl.LightningModule):\n\n    #---->init\n    def __init__(self, model):\n        super(ModelInterface, self).__init__()\n        self.save_hyperparameters()\n        self.load_model()\n        self.criterion = create_loss(CONFIG['base_loss'])\n        \n        self.n_classes = CONFIG['n_classes']\n        self.log_path = CONFIG['log_path']\n\n        #---->acc\n        self.data = [{\"count\": 0, \"correct\": 0} for i in range(self.n_classes)]\n        \n        #---->Metrics\n        if self.n_classes > 2: \n            self.AUROC = torchmetrics.AUROC(num_classes = self.n_classes, average = 'macro', task='multiclass')\n            metrics = torchmetrics.MetricCollection([torchmetrics.Accuracy(num_classes = self.n_classes,\n                                                                           average='micro', task='multiclass'),\n                                                     torchmetrics.CohenKappa(num_classes = self.n_classes, task='multiclass'),\n                                                     torchmetrics.F1Score(num_classes = self.n_classes,\n                                                                     average = 'macro', task='multiclass'),\n                                                     torchmetrics.Recall(average = 'macro',\n                                                                         num_classes = self.n_classes, task='multiclass'),\n                                                     torchmetrics.Precision(average = 'macro',\n                                                                            num_classes = self.n_classes, task='multiclass'),\n                                                     torchmetrics.Specificity(average = 'macro',\n                                                                            num_classes = self.n_classes, task='multiclass')])\n        else : \n            self.AUROC = torchmetrics.AUROC(num_classes=2, average = 'macro', task='multiclass')\n            metrics = torchmetrics.MetricCollection([torchmetrics.Accuracy(num_classes = 2,\n                                                                           average = 'micro', task='multiclass'),\n                                                     torchmetrics.CohenKappa(num_classes = 2, task='multiclass'),\n                                                     torchmetrics.F1Score(num_classes = 2,\n                                                                     average = 'macro', task='multiclass'),\n                                                     torchmetrics.Recall(average = 'macro',\n                                                                         num_classes = 2, task='multiclass'),\n                                                     torchmetrics.Precision(average = 'macro',\n                                                                            num_classes = 2, task='multiclass')])\n        self.balance_accuracy_score = sklearn.metrics.balanced_accuracy_score\n        self.valid_metrics = metrics.clone(prefix = 'val_')\n        self.test_metrics = metrics.clone(prefix = 'test_')\n\n        #--->random\n        self.shuffle = CONFIG['data_shuffle']\n        self.count = 0\n\n        self.validation_step_outputs = []\n        self.test_step_outputs = []\n\n        \n    def forward(self, feats):\n        results_dict  = self.model(feats, mask_ratio=CONFIG['mask_ratio'])\n        return results_dict\n    \n        \n        \n    def configure_optimizers(self):\n        optimizer = optim.RAdam(self.parameters(), lr=CONFIG['lr'], weight_decay=CONFIG['weight_decay'])\n        scheduler = optim.lr_scheduler.CyclicLR(\n            optimizer, base_lr=CONFIG['lr']*0.5, max_lr=CONFIG['lr'] * 1.5,\n            step_size_up=5, cycle_momentum=False, mode=\"triangular2\", verbose=True\n        )\n        return [optimizer], [scheduler]\n    \n\n    #---->remove v_num\n    def get_progress_bar_dict(self):\n        # don't show the version number\n        items = super().get_progress_bar_dict()\n        items.pop(\"v_num\", None)\n        return items\n\n    \n    def training_step(self, batch, batch_idx):\n        \"\"\"每一步的训练/一个epoch中的 training_step\"\"\"\n        #---->inference\n        data, label = batch\n        results_dict = self.forward(data) # 训练时模型返回三个参数\n        logits = results_dict['logits']\n        Y_prob = results_dict['Y_prob']\n        Y_hat = results_dict['Y_hat']\n\n        #---->loss: 袋损失\n        loss = self.criterion(logits, label)\n          \n        #---->acc log\n        Y = int(label)\n        self.data[Y][\"count\"] += 1\n        self.data[Y][\"correct\"] += (Y_hat.item() == Y)\n        \n        return {'loss': loss} \n\n    \n    def on_validation_epoch_start(self):\n        \"\"\"\n            training_epoch 之后，开始 validation 之前进行的操作\n        \"\"\"\n        \n        print(f\"********************************** Epoch:{self.current_epoch} **********************************\")\n        print(\"Training:\")\n        for c in range(self.n_classes):\n            count = self.data[c][\"count\"]\n            correct = self.data[c][\"correct\"]\n            if count == 0: \n                acc = None\n            else:\n                acc = float(correct) / count\n            print('class {}({}): acc {}, correct {}/{}'.format(label_decoder[c], c, acc, correct, count))\n        self.data = [{\"count\": 0, \"correct\": 0} for i in range(self.n_classes)]\n        print()\n\n        \n    def validation_step(self, batch, batch_idx):\n        \"\"\"每一步的验证/一个epoch中的 valid_step\"\"\"\n        data, label = batch\n        results_dict = self(data) # 验证时模型返回三个参数\n        logits = results_dict['logits']\n        Y_prob = results_dict['Y_prob']\n        Y_hat = results_dict['Y_hat']\n        \n        #---->acc log\n        Y = int(label)\n        self.data[Y][\"count\"] += 1\n        self.data[Y][\"correct\"] += (Y_hat.item() == Y)\n        \n        self.validation_step_outputs.append({'logits' : logits, 'Y_prob' : Y_prob, 'Y_hat' : Y_hat, 'label' : label})\n        return {'logits' : logits, 'Y_prob' : Y_prob, 'Y_hat' : Y_hat, 'label' : label}\n\n\n    def on_validation_epoch_end(self):\n        \"\"\"\n            vaild_step 之后进行的操作\n        \"\"\"\n        logits = torch.cat([x['logits'] for x in self.validation_step_outputs], dim = 0) # [45, 5]\n        probs = torch.cat([x['Y_prob'] for x in self.validation_step_outputs], dim = 0) # [45, 5]\n        max_probs = torch.stack([x['Y_hat'] for x in self.validation_step_outputs])\n        target = torch.stack([x['label'] for x in self.validation_step_outputs], dim = 0)\n\n        #---->记录指标\n        self.log('bal_acc', self.balance_accuracy_score(max_probs.cpu(), target.cpu()), prog_bar=True, on_epoch=True, logger=True)\n        self.log('val_loss', cross_entropy_torch(logits, target), prog_bar=True, on_epoch=True, logger=True)\n        self.log('auc', self.AUROC(probs, target.squeeze()), prog_bar=True, on_epoch=True, logger=True)\n        self.log_dict(self.valid_metrics(max_probs.squeeze(0) , target.squeeze(0)),\n                          on_epoch = True, logger = True)\n\n        #---->acc log\n        print(\"Validation:\")\n        for c in range(self.n_classes):\n            count = self.data[c][\"count\"]\n            correct = self.data[c][\"correct\"]\n            if count == 0: \n                acc = None\n            else:\n                acc = float(correct) / count\n            print('class {}({}): acc {}, correct {}/{}'.format(label_decoder[c], c, acc, correct, count))\n        self.data = [{\"count\": 0, \"correct\": 0} for i in range(self.n_classes)]\n        \n        #---->random, if shuffle data, change seed\n        if self.shuffle == True:\n            self.count = self.count+1\n            random.seed(self.count*50)\n            \n        print()\n        self.validation_step_outputs.clear()\n    \n    \n    def test_step(self, batch, batch_idx):\n        \"\"\"测试时才会进行的 test_step\"\"\"\n        data, label = batch\n        results_dict = self(data)\n        logits = results_dict['logits']\n        Y_prob = results_dict['Y_prob']\n        Y_hat = results_dict['Y_hat']\n\n        #---->acc log\n        Y = int(label)\n        self.data[Y][\"count\"] += 1\n        self.data[Y][\"correct\"] += (Y_hat.item() == Y)\n        \n        self.test_step_outputs.append({'logits' : logits, 'Y_prob' : Y_prob, 'Y_hat' : Y_hat, 'label' : label})\n        return {'logits' : logits, 'Y_prob' : Y_prob, 'Y_hat' : Y_hat, 'label' : label}\n\n    \n    def on_test_epoch_end(self):\n        \"\"\"\n            test_step 之后会进行的操作\n        \"\"\"\n        probs = torch.cat([x['Y_prob'] for x in self.test_step_outputs], dim = 0)\n        max_probs = torch.stack([x['Y_hat'] for x in self.test_step_outputs])\n        target = torch.stack([x['label'] for x in self.test_step_outputs], dim = 0)\n        \n        #---->\n        bal_acc = self.balance_accuracy_score(max_probs.cpu(), target.cpu())\n        auc = self.AUROC(probs, target.squeeze())\n        metrics = self.test_metrics(max_probs.squeeze() , target.squeeze())\n        metrics['auc'] = auc\n        metrics['bal_acc'] = bal_acc\n        \n        for keys, values in metrics.items():\n            print(f'{keys} = {values}')\n            metrics[keys] = values\n            \n        #---->acc log\n        print(\"Test:\")\n        for c in range(self.n_classes):\n            count = self.data[c][\"count\"]\n            correct = self.data[c][\"correct\"]\n            if count == 0: \n                acc = None\n            else:\n                acc = float(correct) / count\n            print('class {}({}): acc {}, correct {}/{}'.format(label_decoder[c], c, acc, correct, count))\n        self.data = [{\"count\": 0, \"correct\": 0} for i in range(self.n_classes)]\n        #---->\n        result = pd.DataFrame([metrics])\n        os.makedirs(os.path.join(self.log_path, f\"fold{CONFIG['fold']}\"), exist_ok=True)\n        result.to_csv(os.path.join(self.log_path, f\"fold{CONFIG['fold']}\", f\"result_{CONFIG['model_path']}.csv\"))\n        print()\n        self.test_step_outputs.clear()\n\n\n    def load_model(self):\n        try:\n            Model = MultipleMILTransformer\n        except:\n            raise ValueError('Invalid Module File Name or Invalid Class Name!')\n        self.model = self.instancialize(Model)\n        pass\n\n    \n    def instancialize(self, Model, **other_args):\n        \"\"\" Instancialize a model using the corresponding parameters\n            from self.hparams dictionary. You can also input any args\n            to overwrite the corresponding value in self.hparams.\n        \"\"\"\n        class_args = inspect.getargspec(Model.__init__).args[1:]\n        inkeys = self.hparams.model.keys()\n        args1 = {}\n        for arg in class_args:\n            if arg in inkeys:\n                args1[arg] = self.hparams.model[arg]\n        args1.update(other_args)\n        return Model(**args1)","metadata":{"execution":{"iopub.status.busy":"2024-01-03T11:14:33.413032Z","iopub.execute_input":"2024-01-03T11:14:33.413393Z","iopub.status.idle":"2024-01-03T11:14:33.460171Z","shell.execute_reply.started":"2024-01-03T11:14:33.413365Z","shell.execute_reply":"2024-01-03T11:14:33.458984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Initialize seed\npl.seed_everything(CONFIG['seed'])\n\nmetrics_list = []\nfor fold in range(CONFIG['nfold']):\n    CONFIG['fold'] = fold\n    print()\n    print(f\"**************************************** Train fold{CONFIG['fold']} model ****************************************\")\n    \n    # load loggers\n    CONFIG['load_loggers'] = load_loggers()\n\n    # load callbacks\n    CONFIG['callbacks'] = load_callbacks()\n\n    # Define Data\n    DataInterface_dict = {\n                            'train_batch_size': CONFIG['train_dataloader']['batch_size'],\n                            'train_num_workers': CONFIG['train_dataloader']['num_workers'],\n                            'test_batch_size': CONFIG['test_dataloader']['batch_size'],\n                            'test_num_workers': CONFIG['test_dataloader']['num_workers'],\n                            'dataset_name': CONFIG['dataset_name'],\n                        }\n\n    dm = DataInterface(**DataInterface_dict)\n    dm.setup(stage='fit')\n    # Define Model\n    model_dict = {\n        \"name\": CONFIG['name'],\n        \"n_classes\": CONFIG['n_classes'],\n    }\n    model = ModelInterface(model_dict)\n\n    trainer = Trainer(\n            num_sanity_val_steps = 0,\n            accelerator = CONFIG['device'],\n            devices = CONFIG['devices'],\n            logger = CONFIG['load_loggers'],\n            callbacks = CONFIG['callbacks'],\n            max_epochs = CONFIG['epochs'], \n            precision = CONFIG['precision'],  \n            accumulate_grad_batches = CONFIG['grad_acc'],\n#             fast_dev_run=4,\n#             limit_train_batches=0.1,\n#             limit_val_batches=0.1,\n        )\n\n    # 训练\n    trainer.fit(model = model, datamodule = dm)\n    \n    metrics = pd.read_csv(f\"{CONFIG['fold_log_path']}/metrics.csv\")\n    metrics = metrics.sort_values(by='val_loss', ascending=True)\n    del metrics['step']\n    metrics_list.append(metrics.set_index('epoch').head(CONFIG['save_top_k']))\n    display(metrics_list[-1].style.format(precision=4).background_gradient())\n    \n    \n    # 测试\n    model_paths = list(CONFIG['fold_log_path'].glob('*.ckpt'))\n    model_paths = [str(model_path) for model_path in model_paths if 'epoch' in str(model_path)]\n    print(f\"-------> fold{CONFIG['fold']} start test:\")\n    for idx, path in enumerate(model_paths):\n        print(f\"**************************************** Test model idx:{idx} ****************************************\")\n        print(path)\n        CONFIG['model_path'] = os.path.splitext(os.path.basename(path))[0]\n        new_model = ModelInterface(model_dict)\n        ckpt = torch.load(path)\n        new_model.load_state_dict(ckpt['state_dict'])\n        trainer.test(model=new_model, datamodule=dm)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-03T11:14:34.595275Z","iopub.execute_input":"2024-01-03T11:14:34.595642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics_list = [metric.reset_index(drop=True) for metric in metrics_list]\n\nfor metric in metrics_list:\n    display(metric.head(CONFIG['save_top_k']).style.format(precision=4).background_gradient())\n\nfinal_metric = pd.concat(metrics_list, axis=0)\nfinal_metric.mean(0)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}