{"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":7315511,"sourceType":"datasetVersion","datasetId":4245149}],"dockerImageVersionId":30626,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install nystrom-attention\n!pip install pytorch_toolbelt","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-01T10:28:51.200223Z","iopub.execute_input":"2024-01-01T10:28:51.200561Z","iopub.status.idle":"2024-01-01T10:29:17.031329Z","shell.execute_reply.started":"2024-01-01T10:28:51.200535Z","shell.execute_reply":"2024-01-01T10:29:17.030181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!jupyter notebook --NotebookApp.iopub_data_rate_limit=1.0e14","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-01-01T10:29:17.033516Z","iopub.execute_input":"2024-01-01T10:29:17.033834Z","iopub.status.idle":"2024-01-01T10:29:27.713981Z","shell.execute_reply.started":"2024-01-01T10:29:17.033806Z","shell.execute_reply":"2024-01-01T10:29:27.712844Z"},"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 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-01T10:29:27.715384Z","iopub.execute_input":"2024-01-01T10:29:27.715676Z","iopub.status.idle":"2024-01-01T10:29:34.816407Z","shell.execute_reply.started":"2024-01-01T10:29:27.71565Z","shell.execute_reply":"2024-01-01T10:29:34.815411Z"},"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/feature-0.25scale/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\": \"TransMIL\",\n    \"n_classes\": 5,\n    \n    # Optimizer\n    \"opt\": \"lookahead_radam\",\n    \"lr\": 0.0002,\n    \"opt_eps\": None, \n    \"opt_betas\": None,\n    \"momentum\": None, \n    \"weight_decay\": 0.00001,\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-01T10:29:34.819004Z","iopub.execute_input":"2024-01-01T10:29:34.819548Z","iopub.status.idle":"2024-01-01T10:29:34.82947Z","shell.execute_reply.started":"2024-01-01T10:29:34.819513Z","shell.execute_reply":"2024-01-01T10:29:34.828477Z"},"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    if CONFIG['nfold'] >= 2:\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    else:\n        raise Exception('nfold should be bigger than 2.')\n    \ngenerate_dataset_csv(use_tma = True)","metadata":{"execution":{"iopub.status.busy":"2024-01-01T10:29:34.831246Z","iopub.execute_input":"2024-01-01T10:29:34.832349Z","iopub.status.idle":"2024-01-01T10:29:35.070446Z","shell.execute_reply.started":"2024-01-01T10:29:34.832283Z","shell.execute_reply":"2024-01-01T10:29:35.069457Z"},"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-01T10:29:35.071874Z","iopub.execute_input":"2024-01-01T10:29:35.072491Z","iopub.status.idle":"2024-01-01T10:29:35.090681Z","shell.execute_reply.started":"2024-01-01T10:29:35.072454Z","shell.execute_reply":"2024-01-01T10:29:35.089585Z"},"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\n","metadata":{"execution":{"iopub.status.busy":"2024-01-01T10:29:35.092211Z","iopub.execute_input":"2024-01-01T10:29:35.09254Z","iopub.status.idle":"2024-01-01T10:29:35.111296Z","shell.execute_reply.started":"2024-01-01T10:29:35.092491Z","shell.execute_reply":"2024-01-01T10:29:35.11032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass 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    \n# my_dataset = UBCData(state='train')\n# for i in range(5):\n#     feature, label = my_dataset[i]\n#     print(feature.shape, label)","metadata":{"execution":{"iopub.status.busy":"2024-01-01T10:29:35.112377Z","iopub.execute_input":"2024-01-01T10:29:35.112654Z","iopub.status.idle":"2024-01-01T10:29:35.127382Z","shell.execute_reply.started":"2024-01-01T10:29:35.11263Z","shell.execute_reply":"2024-01-01T10:29:35.126341Z"},"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-01T10:29:35.128967Z","iopub.execute_input":"2024-01-01T10:29:35.129351Z","iopub.status.idle":"2024-01-01T10:29:35.484136Z","shell.execute_reply.started":"2024-01-01T10:29:35.129324Z","shell.execute_reply":"2024-01-01T10:29:35.482998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TransLayer(nn.Module):\n\n    def __init__(self, norm_layer=nn.LayerNorm, dim=512):\n        super().__init__()\n        self.norm = norm_layer(dim)\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\n    def forward(self, x):\n        x = x + self.attn(self.norm(x))\n\n        return x\n\n\nclass PPEG(nn.Module):\n    def __init__(self, dim=512):\n        super(PPEG, self).__init__()\n        self.proj = nn.Conv2d(dim, dim, 7, 1, 7//2, groups=dim)\n        self.proj1 = nn.Conv2d(dim, dim, 5, 1, 5//2, groups=dim)\n        self.proj2 = nn.Conv2d(dim, dim, 3, 1, 3//2, groups=dim)\n\n    def forward(self, x, H, W):\n        B, _, C = x.shape\n        cls_token, feat_token = x[:, 0], x[:, 1:]\n        cnn_feat = feat_token.transpose(1, 2).view(B, C, H, W)\n        x = self.proj(cnn_feat)+cnn_feat+self.proj1(cnn_feat)+self.proj2(cnn_feat)\n        x = x.flatten(2).transpose(1, 2)\n        x = torch.cat((cls_token.unsqueeze(1), x), dim=1)\n        return x\n\n\nclass TransMIL(nn.Module):\n    def __init__(self, n_classes):\n        super(TransMIL, self).__init__()\n        self.pos_layer = PPEG(dim=512)\n        self._fc1 = nn.Sequential(nn.Linear(384, 512), nn.ReLU()) # 将输入特征进行升维\n        self.cls_token = nn.Parameter(torch.randn(1, 1, 512))\n        self.n_classes = n_classes\n        self.layer1 = TransLayer(dim=512)\n        self.layer2 = TransLayer(dim=512)\n        self.norm = nn.LayerNorm(512)\n        self._fc2 = nn.Linear(512, self.n_classes)\n\n\n    def forward(self, **kwargs):\n        # 输入特征 H\n        h = kwargs['data'].float() #[B, n, 384]\n        \n        h = self._fc1(h) #[B, n, 512]\n        \n        #---->pad\n        H = h.shape[1]\n        _H, _W = int(np.ceil(np.sqrt(H))), int(np.ceil(np.sqrt(H)))\n        add_length = _H * _W - H\n        h = torch.cat([h, h[:,:add_length,:]],dim = 1) #[B, N, 512]\n\n        #---->cls_token\n        B = h.shape[0]\n        cls_tokens = self.cls_token.expand(B, -1, -1).cuda()\n        h = torch.cat((cls_tokens, h), dim=1)\n\n        #---->Translayer x1\n        h = self.layer1(h) #[B, N, 512]\n\n        #---->PPEG\n        h = self.pos_layer(h, _H, _W) #[B, N, 512]\n        \n        #---->Translayer x2\n        h = self.layer2(h) #[B, N, 512]\n\n        #---->cls_token\n        h = self.norm(h)[:,0]\n\n        #---->predict\n        logits = self._fc2(h) #[B, n_classes] 从这里这届预测出每个包在 n_classes 个类别上的特征映射结果\n        Y_hat = torch.argmax(logits, dim=1) # 选择概率最大的类的下标index作为包的类别\n        Y_prob = F.softmax(logits, dim = 1) # 从第1个维度开始，对每一个包的特征映射进行归一化\n        results_dict = {'logits': logits, 'Y_prob': Y_prob, 'Y_hat': Y_hat} \n        \n        # 返回三个结果，特征映射的结果，归一化的概率，预测出的包的类别，通过标签训练模型\n        return results_dict","metadata":{"execution":{"iopub.status.busy":"2024-01-01T10:29:35.488823Z","iopub.execute_input":"2024-01-01T10:29:35.489377Z","iopub.status.idle":"2024-01-01T10:29:35.517756Z","shell.execute_reply.started":"2024-01-01T10:29:35.489335Z","shell.execute_reply":"2024-01-01T10:29:35.516584Z"},"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.loss = create_loss(CONFIG['base_loss'])\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(data=feats)\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.loss(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 = TransMIL\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-01T10:29:35.519343Z","iopub.execute_input":"2024-01-01T10:29:35.519753Z","iopub.status.idle":"2024-01-01T10:29:35.576232Z","shell.execute_reply.started":"2024-01-01T10:29:35.519718Z","shell.execute_reply":"2024-01-01T10:29:35.575169Z"},"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)\n        ","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-01T10:29:35.577634Z","iopub.execute_input":"2024-01-01T10:29:35.57799Z","iopub.status.idle":"2024-01-01T10:39:55.113619Z","shell.execute_reply.started":"2024-01-01T10:29:35.577959Z","shell.execute_reply":"2024-01-01T10:39:55.112551Z"},"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":{"execution":{"iopub.status.busy":"2024-01-01T10:56:12.379473Z","iopub.execute_input":"2024-01-01T10:56:12.380284Z","iopub.status.idle":"2024-01-01T10:56:12.444113Z","shell.execute_reply.started":"2024-01-01T10:56:12.380245Z","shell.execute_reply":"2024-01-01T10:56:12.443088Z"},"trusted":true},"execution_count":null,"outputs":[]}]}