{"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":6774400,"sourceType":"datasetVersion","datasetId":3895136},{"sourceId":6774553,"sourceType":"datasetVersion","datasetId":3898019},{"sourceId":6981947,"sourceType":"datasetVersion","datasetId":4012490},{"sourceId":6981956,"sourceType":"datasetVersion","datasetId":4012497},{"sourceId":6987951,"sourceType":"datasetVersion","datasetId":4016160},{"sourceId":6995549,"sourceType":"datasetVersion","datasetId":4021094},{"sourceId":147718027,"sourceType":"kernelVersion"}],"dockerImageVersionId":30587,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"import torch\nfrom torchvision.models import resnet50\n\nfrom vit_pytorch.distill import DistillableViT, DistillWrapper\n\nteacher = resnet50(pretrained = True)\n\nv = DistillableViT(\n    image_size = 256,\n    patch_size = 32,\n    num_classes = 1000,\n    dim = 1024,\n    depth = 6,\n    heads = 8,\n    mlp_dim = 2048,\n    dropout = 0.1,\n    emb_dropout = 0.1\n)\n\ndistiller = DistillWrapper(\n    student = v,\n    teacher = teacher,\n    temperature = 3,           # temperature of distillation\n    alpha = 0.5,               # trade between main loss and distillation loss\n    hard = False               # whether to use soft or hard distillation\n)\n\nimg = torch.randn(2, 3, 256, 256)\nlabels = torch.randint(0, 1000, (2,))\n\nloss = distiller(img, labels)\nloss.backward()\n# Cancer🔬 Classification: Baseline with ⚡`lightning`\n\n**It is continuation of EDA: https://www.kaggle.com/code/jirkaborovec/cancer-subtype-explore-data-images**","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"import os, glob\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nDATASET_FOLDER = \"/kaggle/input/UBC-OCEAN/\"\nDATASET_IMAGES = \"/kaggle/input/cancer-subtype-eda-load-wsi-prune-bg/train_thumbnails/\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"_kg_hide-input":true,"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:04:45.959263Z","iopub.execute_input":"2023-11-18T11:04:45.959666Z","iopub.status.idle":"2023-11-18T11:04:45.965125Z","shell.execute_reply.started":"2023-11-18T11:04:45.959629Z","shell.execute_reply":"2023-11-18T11:04:45.963847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Checkout some labels","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"df_train = pd.read_csv(os.path.join(DATASET_FOLDER, \"train.csv\"))\n# labels = list(df_train[\"label\"].unique())\nprint(f\"Dataset/train size: {len(df_train)}\")\ndisplay(df_train.head())","metadata":{"_kg_hide-input":true,"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:04:45.967224Z","iopub.execute_input":"2023-11-18T11:04:45.967535Z","iopub.status.idle":"2023-11-18T11:04:45.991039Z","shell.execute_reply.started":"2023-11-18T11:04:45.967509Z","shell.execute_reply":"2023-11-18T11:04:45.990077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_= df_train[[\"label\"]].value_counts().plot.pie(autopct='%1.1f%%', ylabel=\"label\", figsize=(3,3))","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:04:45.99268Z","iopub.execute_input":"2023-11-18T11:04:45.992995Z","iopub.status.idle":"2023-11-18T11:04:46.136858Z","shell.execute_reply.started":"2023-11-18T11:04:45.992969Z","shell.execute_reply":"2023-11-18T11:04:46.135595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Show some samples 🖼️ per class\n\nNote that not all images has thumbnails","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"df_train[\"path_thumbnail\"] = df_train['image_id'].apply(lambda id: f\"{id}_thumbnail.png\")\ndf_train[\"thumbnail_exists\"] = df_train['path_thumbnail'].apply(\n    lambda pth: os.path.isfile(os.path.join(DATASET_IMAGES, pth)))\ndisplay(df_train.head())\ndf_train = df_train[df_train[\"thumbnail_exists\"] == True]\nprint(f\"size: {len(df_train)}\")","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:04:46.138631Z","iopub.execute_input":"2023-11-18T11:04:46.139319Z","iopub.status.idle":"2023-11-18T11:04:46.379644Z","shell.execute_reply.started":"2023-11-18T11:04:46.139273Z","shell.execute_reply":"2023-11-18T11:04:46.378729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nnb_samples = 6\nn, m = len(np.unique(df_train['label'])), nb_samples,\nfig, axarr = plt.subplots(nrows=n, ncols=m, figsize=(m * 2, n * 2))\nfor ilb, (lb, df_) in enumerate(df_train.groupby('label')):\n    img_ids = list(df_['image_id'])\n    for i in range(m):\n        img_path = os.path.join(DATASET_IMAGES, f\"{img_ids[i]}_thumbnail.png\")\n        img = plt.imread(img_path)\n        axarr[ilb, i].imshow(img)\n        if i == 0:\n            axarr[ilb, i].set_title(f\"{lb} #{len(df_)}\")\n        axarr[ilb, i].set_xticks([])\n        axarr[ilb, i].set_yticks([])\n_= plt.axis('off')","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:04:46.382536Z","iopub.execute_input":"2023-11-18T11:04:46.382839Z","iopub.status.idle":"2023-11-18T11:05:18.545586Z","shell.execute_reply.started":"2023-11-18T11:04:46.382812Z","shell.execute_reply":"2023-11-18T11:05:18.544686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data pre-processing\n\n### Sacling down images\n\nWe could not fit this huge image size to any NEt so just as offline process lets downscale it to about 1024x1024","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"def prune_image_rows_cols(img, thr=0.001):\n    # delete empty columns\n    for l in reversed(range(img.shape[1])):\n        if (np.sum(img[:, l]) / float(img.shape[0])) < thr:\n            img = np.delete(img, l, 1)\n    # delete empty rows\n    for l in reversed(range(img.shape[0])):\n        if (np.sum(img[l, :]) / float(img.shape[1])) < thr:\n            img = np.delete(img, l, 0)\n    return img","metadata":{"_kg_hide-input":true,"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:05:18.546848Z","iopub.execute_input":"2023-11-18T11:05:18.547135Z","iopub.status.idle":"2023-11-18T11:05:18.553847Z","shell.execute_reply.started":"2023-11-18T11:05:18.547109Z","shell.execute_reply":"2023-11-18T11:05:18.553033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prune_scale_image(img_path: str, out_dir: str, size: int = 1024) -> None:\n    img = np.array(Image.open(img_path))\n    img = prune_image_rows_cols(img)\n    mask = np.sum(img[..., :3], axis=2) == 0\n    img[mask, :] = 255\n    img = Image.fromarray(img)\n    img.thumbnail((size, size))\n    img.save(os.path.join(out_dir, os.path.basename(img_path)))","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:05:18.555035Z","iopub.execute_input":"2023-11-18T11:05:18.555307Z","iopub.status.idle":"2023-11-18T11:05:18.568234Z","shell.execute_reply.started":"2023-11-18T11:05:18.555283Z","shell.execute_reply":"2023-11-18T11:05:18.567304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nfrom joblib import Parallel, delayed\n\n! mkdir -p train_thumbnails\n! rm -f train_thumbnails/*.png\n\nls = glob.glob(os.path.join(DATASET_IMAGES, '*.png'))\nprint(f\"found images: {len(ls)}\")\n\n# for p_img in tqdm(ls):\n#     prune_scale_image(p_img, \"./train_thumbnails\")\n    \n_= Parallel(n_jobs=4)(\n    delayed(prune_scale_image)(p_img, \"./train_thumbnails\") for p_img in tqdm(ls)\n)\nls = glob.glob(os.path.join(\"./train_thumbnails\", '*.png'))\nprint(f\"found images: {len(ls)}\")","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:05:18.569572Z","iopub.execute_input":"2023-11-18T11:05:18.569858Z","iopub.status.idle":"2023-11-18T11:08:32.038368Z","shell.execute_reply.started":"2023-11-18T11:05:18.569834Z","shell.execute_reply":"2023-11-18T11:08:32.037298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Color 🦩 normalizations","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"def _color_means(img_path):\n    img = np.array(Image.open(img_path))\n    if np.max(img) > 1.5:\n        img = img / 255.0\n    clr_mean = {i: np.mean(img[..., i]) for i in range(3)}\n    clr_std = {i: np.std(img[..., i]) for i in range(3)}\n    return clr_mean, clr_std\n\n# os.path.join(DATASET_SMALL_FOLDER, \"train_images\")\nls_images = glob.glob(os.path.join(\"./train_thumbnails\", \"*.png\"))\nclr_mean_std = Parallel(n_jobs=os.cpu_count())(delayed(_color_means)(fn) for fn in tqdm(ls_images))","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:08:32.039973Z","iopub.execute_input":"2023-11-18T11:08:32.040294Z","iopub.status.idle":"2023-11-18T11:08:43.580977Z","shell.execute_reply.started":"2023-11-18T11:08:32.040266Z","shell.execute_reply":"2023-11-18T11:08:43.580137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_color_mean = pd.DataFrame([c[0] for c in clr_mean_std]).describe()\ndisplay(img_color_mean.T)\nimg_color_std = pd.DataFrame([c[1] for c in clr_mean_std]).describe()\ndisplay(img_color_std.T)\n\nimg_color_mean = list(img_color_mean.T[\"mean\"])\nimg_color_std = list(img_color_std.T[\"mean\"])\nprint(f\"{img_color_mean=}\\n{img_color_std=}\")","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:08:43.582162Z","iopub.execute_input":"2023-11-18T11:08:43.582473Z","iopub.status.idle":"2023-11-18T11:08:43.630399Z","shell.execute_reply.started":"2023-11-18T11:08:43.582435Z","shell.execute_reply":"2023-11-18T11:08:43.629409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset & DataModule\n\nCreating standard PyTorch dataset to define how the data shall be loaded and set representations. We define the sample pair as:\n- RGB image\n- one-hot lable encding\n\nA DataModule standardizes the training, val, test splits, data preparation and transforms. The main advantage is consistent data splits, data preparation and transforms across models.","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\n\nclass CancerThumbnailDataset(Dataset):\n    split: float = 0.90\n\n    def __init__(\n        self,\n        df_data,\n        path_img_dir: str =  'train_thumbnails',\n        transforms = None,\n        mode: str = 'train',\n        labels_lut = None\n    ):\n        # 数据集的初始化方法\n        self.path_img_dir = path_img_dir  # 图像文件的路径\n        self.transforms = transforms  # 数据变换\n        self.mode = mode  # 数据集模式（训练或验证）\n\n        self.data = df_data  # 数据集的DataFrame\n        self.labels_unique = sorted(self.data[\"label\"].unique())  # 唯一标签列表\n        self.labels_lut = labels_lut or {lb: i for i, lb in enumerate(self.labels_unique)}  # 标签字典\n        # 随机打乱数据\n        self.data = self.data.sample(frac=1, random_state=42).reset_index(drop=True)\n\n        # 划分数据集\n        assert 0.0 <= self.split <= 1.0\n        frac = int(self.split * len(self.data))\n        self.data = self.data[:frac] if mode == 'train' else self.data[frac:]\n        self.img_names = [f\"{id}_thumbnail.png\" for id in self.data[\"image_id\"]]\n        self.labels = list(self.data['label'])  # 数据集的标签列表\n\n    @property\n    def num_classes(self) -> int:\n        # 返回数据集中的类别数\n        return len(self.labels_lut)\n\n    def to_one_hot(self, label: str) -> int:\n        # 将标签转换为索引\n        return self.labels_lut[label]\n\n    def __getitem__(self, idx: int) -> tuple:\n        # 根据索引获取图像和标签\n        img_path = os.path.join(self.path_img_dir, self.img_names[idx])\n        assert os.path.isfile(img_path), f\"missing: {img_path}\"\n        img = plt.imread(img_path)[..., :3]  # 读取图像\n        if np.max(img) < 1.5:\n            img = np.clip(img * 255, 0, 255).astype(np.uint8)\n        labels = self.to_one_hot(self.labels[idx])  # 转换标签为独热编码\n\n        # 数据增强\n        if self.transforms:\n            img = self.transforms(Image.fromarray(img))\n\n        return img, torch.tensor(labels).to(int)\n\n    def __len__(self) -> int:\n        # 返回数据集的长度\n        return len(self.data)\n\n# ==============================\n# ==============================\n\n# 创建数据集实例\ndataset = CancerThumbnailDataset(df_train)\n\n# 可视化数据集的前9个样本\nfig = plt.figure(figsize=(10, 8))\nfor i in range(9):\n    img, lb = dataset[i]\n    ax = fig.add_subplot(3, 3, i + 1, xticks=[], yticks=[])\n    ax.imshow(img)\n    ax.set_title(f\"{lb}\\n dims: {img.shape} /{np.max(img)}\")\n","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:08:43.632022Z","iopub.execute_input":"2023-11-18T11:08:43.632314Z","iopub.status.idle":"2023-11-18T11:08:49.40038Z","shell.execute_reply.started":"2023-11-18T11:08:43.632289Z","shell.execute_reply":"2023-11-18T11:08:49.399475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Let us define some standard image augmentaion procedures and color normalizations...","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"from torchvision import transforms as T\nfrom torchvision.transforms import InterpolationMode\n\nTRAIN_TRANSFORM = T.Compose([\n    T.RandomRotation(45, fill=255),\n    T.RandomCrop(512, pad_if_needed=True, padding_mode=\"reflect\"),\n    #T.RandomResizedCrop(512, interpolation=InterpolationMode.BICUBIC, antialias=True),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.ToTensor(),\n    #T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n    T.Normalize(img_color_mean, img_color_std),  # custom\n])\n\nVALID_TRANSFORM = T.Compose([\n    T.CenterCrop(512),\n    T.ToTensor(),\n    #T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n    T.Normalize(img_color_mean, img_color_std),  # custom\n])","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:08:49.401802Z","iopub.execute_input":"2023-11-18T11:08:49.40221Z","iopub.status.idle":"2023-11-18T11:08:49.723288Z","shell.execute_reply.started":"2023-11-18T11:08:49.402184Z","shell.execute_reply":"2023-11-18T11:08:49.722506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### The DataModule include creating training and validation dataset with given split and feading it to particular data loaders...","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"import multiprocessing as mproc\nimport pytorch_lightning as pl\nfrom torch.utils.data import DataLoader\n\n# 定义一个PyTorch Lightning数据模块\nclass CancerSubtypeDM(pl.LightningDataModule):\n\n    def __init__(\n        self,\n        df_data,\n        path_img_dir: str = 'train_thumbnails',\n        batch_size: int = 32,\n        num_workers: int = None,\n        train_transforms = TRAIN_TRANSFORM,\n        valid_transforms = VALID_TRANSFORM\n    ):\n        super().__init__()\n        self.df_data = df_data\n        self.path_img_dir = path_img_dir\n        self.batch_size = batch_size\n        # 确定使用的工作进程数（默认为CPU核心数）\n        self.num_workers = num_workers or mproc.cpu_count()\n        self.train_dataset = None\n        self.valid_dataset = None\n        self.train_transforms = train_transforms\n        self.valid_transforms = valid_transforms\n\n    # 准备数据的方法，可以留空\n    def prepare_data(self):\n        pass\n\n    # 返回类别数\n    @property\n    def num_classes(self) -> int:\n        assert self.train_dataset and self.valid_dataset\n        return len(set(self.train_dataset.labels_unique + self.valid_dataset.labels_unique))\n\n    # 设置训练和验证数据集\n    def setup(self, stage=None):\n        self.train_dataset = CancerThumbnailDataset(\n            self.df_data, self.path_img_dir, mode='train', transforms=self.train_transforms)\n        print(f\"training dataset: {len(self.train_dataset)}\")\n        self.valid_dataset = CancerThumbnailDataset(\n            self.df_data, self.path_img_dir, mode='valid', transforms=self.valid_transforms,\n            labels_lut=self.train_dataset.labels_lut)\n        print(f\"validation dataset: {len(self.valid_dataset)}\")\n\n    # 创建训练数据加载器\n    def train_dataloader(self):\n        return DataLoader(\n            self.train_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=True,\n        )\n\n    # 创建验证数据加载器\n    def val_dataloader(self):\n        return DataLoader(\n            self.valid_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=False,\n        )\n\n    # 创建测试数据加载器\n    def test_dataloader(self):\n        pass\n\n# 创建数据模块对象\ndm = CancerSubtypeDM(df_train, batch_size=8)\n# 设置数据模块\ndm.setup()\n# 打印类别数\nprint(dm.num_classes)\n\n# 快速查看数据集中的一批图像\nfig = plt.figure(figsize=(3, 7))\nfor imgs, lbs in dm.train_dataloader():\n    print(f'batch labels: {torch.sum(lbs, axis=0)}')\n    print(f'image size: {imgs[0].shape}')\n    for i in range(3):\n        ax = fig.add_subplot(3, 1, i + 1, xticks=[], yticks=[])\n        ax.imshow(np.rollaxis(imgs[i].numpy(), 0, 3))\n        ax.set_title(lbs[i])\n    break\n","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:08:49.727435Z","iopub.execute_input":"2023-11-18T11:08:49.727764Z","iopub.status.idle":"2023-11-18T11:08:55.039666Z","shell.execute_reply.started":"2023-11-18T11:08:49.727738Z","shell.execute_reply":"2023-11-18T11:08:55.038634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ViT Model\n\nWe start with some standard ViT models.\nThen we define Ligthning module including training and validation step and configure optimizer/schedular.","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"!pip install -q /kaggle/input/einops/einops-0.7.0-py3-none-any.whl\n!pip install -q /kaggle/input/vit-pytorch/vit_pytorch-1.6.4-py3-none-any.whl","metadata":{"_kg_hide-output":true,"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T12:19:58.125284Z","iopub.execute_input":"2023-11-18T12:19:58.125551Z","iopub.status.idle":"2023-11-18T12:21:02.465856Z","shell.execute_reply.started":"2023-11-18T12:19:58.125525Z","shell.execute_reply":"2023-11-18T12:21:02.464673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from __future__ import print_function\n\nfrom itertools import chain\nimport random\n\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom sklearn.model_selection import train_test_split\nfrom torch.optim.lr_scheduler import StepLR\nfrom tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-11-18T11:09:08.343295Z","iopub.execute_input":"2023-11-18T11:09:08.344266Z","iopub.status.idle":"2023-11-18T11:09:08.819219Z","shell.execute_reply.started":"2023-11-18T11:09:08.344221Z","shell.execute_reply":"2023-11-18T11:09:08.818347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training settings","metadata":{}},{"cell_type":"code","source":"batch_size = 16\nepochs = 20\nlr = 0.01\n# gamma = 0.7","metadata":{"execution":{"iopub.status.busy":"2023-11-18T11:09:08.820504Z","iopub.execute_input":"2023-11-18T11:09:08.820793Z","iopub.status.idle":"2023-11-18T11:09:08.825419Z","shell.execute_reply.started":"2023-11-18T11:09:08.820768Z","shell.execute_reply":"2023-11-18T11:09:08.824427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2023-11-18T11:09:08.826556Z","iopub.execute_input":"2023-11-18T11:09:08.826845Z","iopub.status.idle":"2023-11-18T11:09:08.858342Z","shell.execute_reply.started":"2023-11-18T11:09:08.826819Z","shell.execute_reply":"2023-11-18T11:09:08.857371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visual Transformer","metadata":{}},{"cell_type":"code","source":"from vit_pytorch.vit_for_small_dataset import ViT\n\nmodel = ViT(\n    image_size = 512,\n    patch_size = 16,\n    num_classes = 5,\n    dim = 1024,\n    depth = 6,\n    heads = 16,\n    mlp_dim = 2048,\n    dropout = 0.1,\n    emb_dropout = 0.1\n).to(device)","metadata":{"execution":{"iopub.status.busy":"2023-11-18T11:09:08.859719Z","iopub.execute_input":"2023-11-18T11:09:08.860004Z","iopub.status.idle":"2023-11-18T11:09:12.963427Z","shell.execute_reply.started":"2023-11-18T11:09:08.859972Z","shell.execute_reply":"2023-11-18T11:09:12.96253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"# loss function\ncriterion = nn.CrossEntropyLoss()\n# optimizer\noptimizer = optim.Adam(model.parameters(), lr=lr, weight_decay=1e-4)\n# scheduler\n# scheduler = StepLR(optimizer, step_size=1, gamma=gamma)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=5, verbose=True)","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:09:12.964634Z","iopub.execute_input":"2023-11-18T11:09:12.964937Z","iopub.status.idle":"2023-11-18T11:09:12.970959Z","shell.execute_reply.started":"2023-11-18T11:09:12.964912Z","shell.execute_reply":"2023-11-18T11:09:12.969875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 初始化记录指标的空列表\ntrain_losses = []\ntrain_accuracies = []\nval_losses = []\nval_accuracies = []\n\n# 循环训练\nfor epoch in range(epochs):\n    epoch_loss = 0\n    epoch_accuracy = 0\n\n    for data, label in tqdm(dm.train_dataloader()):\n        data = data.to(device)\n        label = label.to(device)\n\n        output = model(data)\n        loss = criterion(output, label.long())\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        acc = (output.argmax(dim=1) == label).float().mean()\n        epoch_accuracy += acc / len(dm.train_dataloader())\n        epoch_loss += loss / len(dm.train_dataloader())\n\n    # 获取当前学习率\n    current_lr = optimizer.param_groups[0]['lr']\n    \n    # 记录训练集的损失和准确率\n    train_losses.append(epoch_loss.item())\n    train_accuracies.append(epoch_accuracy.item())\n\n    with torch.no_grad():\n        epoch_val_accuracy = 0\n        epoch_val_loss = 0\n        for data, label in dm.val_dataloader():\n            data = data.to(device)\n            label = label.to(device)\n\n            val_output = model(data)\n            val_loss = criterion(val_output, label.long())\n\n            acc = (val_output.argmax(dim=1) == label).float().mean()\n            epoch_val_accuracy += acc / len(dm.val_dataloader())\n            epoch_val_loss += val_loss / len(dm.val_dataloader())\n            \n        # 记录验证集的损失和准确率\n        val_losses.append(epoch_val_loss.item())\n        val_accuracies.append(epoch_val_accuracy.item())\n       \n    # 更新学习率\n    scheduler.step(epoch_val_loss)\n\n    print(\n        f\"Epoch : {epoch+1} - loss : {epoch_loss:.4f} - acc: {epoch_accuracy:.4f} - val_loss : {epoch_val_loss:.4f} - val_acc: {epoch_val_accuracy:.4f} - lr: {current_lr}\\n\"\n    )","metadata":{"_kg_hide-output":true,"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:09:12.972144Z","iopub.execute_input":"2023-11-18T11:09:12.97242Z","iopub.status.idle":"2023-11-18T11:25:16.447412Z","shell.execute_reply.started":"2023-11-18T11:09:12.972395Z","shell.execute_reply":"2023-11-18T11:25:16.446166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Quick visualization of the training process...","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"# 绘制指标关系图\nimport matplotlib.pyplot as plt\n\n# 绘制训练集和验证集的损失\nplt.figure(figsize=(12, 4))\nplt.subplot(1, 2, 1)\nplt.plot(train_losses, label='train loss')\nplt.plot(val_losses, label='val loss')\nplt.title('Training and Validation Loss')\nplt.legend()\n\n# 绘制训练集和验证集的准确率\nplt.subplot(1, 2, 2)\nplt.plot(train_accuracies, label='train accuracy')\nplt.plot(val_accuracies, label='val accuracy')\nplt.title('Training and Validation Accuracy')\nplt.legend()\n\nplt.show()","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:25:16.449108Z","iopub.execute_input":"2023-11-18T11:25:16.449408Z","iopub.status.idle":"2023-11-18T11:25:16.952781Z","shell.execute_reply.started":"2023-11-18T11:25:16.44938Z","shell.execute_reply":"2023-11-18T11:25:16.951844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Save the model!","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"torch.save(model.state_dict(), './trained-vit.pt')","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:25:16.954242Z","iopub.execute_input":"2023-11-18T11:25:16.954881Z","iopub.status.idle":"2023-11-18T11:25:17.259041Z","shell.execute_reply.started":"2023-11-18T11:25:16.954847Z","shell.execute_reply":"2023-11-18T11:25:17.257977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Inference coming in https://www.kaggle.com/code/jirkaborovec/cancer-subtype-lightning-torch-inference","metadata":{"pycharm":{"name":"#%% md\n"}}},{"cell_type":"code","source":"import torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset\n\nclass TilesFolderDataset(Dataset):\n    def __init__(\n        self,\n        folder_tiles: list,  # 将参数名修改为 folder_tiles，表示这是包含图像文件路径的列表\n        image_ext: str = '.png',\n        transforms=None\n    ):\n        self.transforms = transforms\n        self.imgs = folder_tiles  # 直接使用 folder_tiles，因为它已经是图像文件的路径列表\n\n    def __getitem__(self, idx: int) -> tuple:\n        img_path = self.imgs[idx]\n        \n        # 断言图像文件存在\n        assert os.path.isfile(img_path), f\"Missing file: {img_path}\"\n        \n        # 读取图像并去除 alpha 通道\n        img = np.array(Image.open(img_path))[..., :3]\n        \n        # 过滤背景\n        mask = np.sum(img, axis=2) == 0\n        img[mask, :] = 255\n        \n        # 如果图像的最大值小于1.5，则将其缩放到0-255的范围\n        if np.max(img) < 1.5:\n            img = np.clip(img * 255, 0, 255).astype(np.uint8)\n        \n        # 数据增强\n        if self.transforms:\n            img = self.transforms(Image.fromarray(img))\n        \n        return img\n\n    def __len__(self) -> int:\n        return len(self.imgs)\n","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:25:17.26043Z","iopub.execute_input":"2023-11-18T11:25:17.260762Z","iopub.status.idle":"2023-11-18T11:25:17.269948Z","shell.execute_reply.started":"2023-11-18T11:25:17.260736Z","shell.execute_reply":"2023-11-18T11:25:17.269084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 从 CSV 文件中读取测试数据集\ndf_test = pd.read_csv(os.path.join(DATASET_FOLDER, \"test.csv\"))\n\n# 为每个测试样本设置默认标签 'HGSC'\ndf_test['label'] = ['HGSC'] * len(df_test)\n\n# 打印测试数据集的大小（行数）\nprint(f\"Dataset/test size: {len(df_test)}\")\n\n# 显示测试数据集的前几行，以便查看数据结构和内容\ndisplay(df_test.head())\n","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:25:17.271132Z","iopub.execute_input":"2023-11-18T11:25:17.271411Z","iopub.status.idle":"2023-11-18T11:25:17.298641Z","shell.execute_reply.started":"2023-11-18T11:25:17.271387Z","shell.execute_reply":"2023-11-18T11:25:17.297744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cat /kaggle/input/UBC-OCEAN/sample_submission.csv","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:25:17.300046Z","iopub.execute_input":"2023-11-18T11:25:17.300787Z","iopub.status.idle":"2023-11-18T11:25:18.28174Z","shell.execute_reply.started":"2023-11-18T11:25:17.30075Z","shell.execute_reply":"2023-11-18T11:25:18.280676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls /kaggle/input/pyvips-python-and-deb-package-gpu\n# intall the deb packages\n!yes | dpkg -i --force-depends /kaggle/input/pyvips-python-and-deb-package-gpu/linux_packages/archives/*.deb\n# install the python wrapper\n!pip install pyvips -f /kaggle/input/pyvips-python-and-deb-package-gpu/python_packages/ --no-index","metadata":{"_kg_hide-output":false,"_kg_hide-input":true,"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:25:18.283295Z","iopub.execute_input":"2023-11-18T11:25:18.283649Z","iopub.status.idle":"2023-11-18T11:26:19.486268Z","shell.execute_reply.started":"2023-11-18T11:25:18.283602Z","shell.execute_reply":"2023-11-18T11:26:19.485194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pyvips\nimport random\n\nDATASET_FOLDER = \"/kaggle/input/UBC-OCEAN/\"\nIMAGES_FOLDER = \"/kaggle/working/test_tiles\"","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:26:19.488135Z","iopub.execute_input":"2023-11-18T11:26:19.488571Z","iopub.status.idle":"2023-11-18T11:26:19.799016Z","shell.execute_reply.started":"2023-11-18T11:26:19.488532Z","shell.execute_reply":"2023-11-18T11:26:19.797995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 读取训练数据集的 CSV 文件，文件路径由 DATASET_FOLDER 和 \"train.csv\" 构成\ndf_train = pd.read_csv(os.path.join(DATASET_FOLDER, \"train.csv\"))\n\n# 获取训练数据集中 \"label\" 列的唯一值，并按字母顺序排序\nlabels = sorted(df_train[\"label\"].unique())\n\n# 输出所有标签的唯一值列表\nprint(f\"{labels=}\")\n","metadata":{"execution":{"iopub.status.busy":"2023-11-18T11:26:19.800279Z","iopub.execute_input":"2023-11-18T11:26:19.800594Z","iopub.status.idle":"2023-11-18T11:26:19.811392Z","shell.execute_reply.started":"2023-11-18T11:26:19.800567Z","shell.execute_reply":"2023-11-18T11:26:19.810512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def extract_image_tiles(\n    p_img,           # 输入图像的路径\n    folder,          # 存储提取的图像块的文件夹路径\n    size: int = 2048,  # 每个图像块的大小\n    scale: float = 0.5,  # 图像块的缩放比例\n    drop_thr: float = 0.6,  # 阈值，用于判断是否丢弃几乎为空的图像块\n    white_thr: int = 240,  # 白色的阈值，用于判断是否为空白图像块\n    max_samples: int = 50  # 最大提取的图像块数量\n) -> list:\n    # 从文件路径中提取图像的名称和扩展名\n    name, _ = os.path.splitext(os.path.basename(p_img))\n    \n    # 使用 pyvips 库加载图像\n    im = pyvips.Image.new_from_file(p_img)\n    \n    # 设置图像块的宽度和高度\n    w = h = size\n    \n    # 生成图像块的索引\n    idxs = [(y, y + h, x, x + w) for y in range(0, im.height, h) for x in range(0, im.width, w)]\n    \n    # 对索引进行随机子采样\n    max_samples = max_samples if isinstance(max_samples, int) else int(len(idxs) * max_samples)\n    random.shuffle(idxs)\n    \n    # 存储提取的图像块文件路径的列表\n    files = []\n    \n    # 确保文件夹存在或创建文件夹\n    os.makedirs(folder, exist_ok=True)\n    \n    for y, y_, x, x_ in idxs:\n        # 使用 pyvips 库裁剪图像块\n        tile = im.crop(x, y, min(w, im.width - x), min(h, im.height - y)).numpy()[..., :3]\n        \n        # 处理非标准大小的图像块\n        if tile.shape[:2] != (h, w):\n            tile_ = tile\n            tile_size = (h, w) if tile.ndim == 2 else (h, w, tile.shape[2])\n            tile = np.zeros(tile_size, dtype=tile.dtype)\n            tile[:tile_.shape[0], :tile_.shape[1], ...] = tile_\n        \n        # 处理全黑的图像块\n        black_bg = np.sum(tile, axis=2) == 0\n        tile[black_bg, :] = 255\n        \n        # 处理白色背景的图像块\n        mask_bg = np.mean(tile, axis=2) > white_thr\n        if np.sum(mask_bg) >= (np.prod(mask_bg.shape) * drop_thr):\n            # 如果图像块几乎为空，跳过\n            continue\n        \n        # 生成图像块文件的路径\n        p_img_path = os.path.join(folder, f\"{int(x_ / w)}-{int(y_ / h)}.png\")\n        \n        # 将图像块保存为文件\n        new_size = int(size * scale), int(size * scale)\n        Image.fromarray(tile).resize(new_size, Image.LANCZOS).save(p_img_path)\n        \n        # 将文件路径添加到列表中\n        files.append(p_img_path)\n        \n        # 设置计数器检查，以便一些空的图像块可能被提前跳过\n        if len(files) >= max_samples:\n            break\n    \n    # 返回提取的图像块文件路径列表\n    return files\n","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:26:19.813161Z","iopub.execute_input":"2023-11-18T11:26:19.813687Z","iopub.status.idle":"2023-11-18T11:26:19.828541Z","shell.execute_reply.started":"2023-11-18T11:26:19.813651Z","shell.execute_reply":"2023-11-18T11:26:19.827538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\nfrom torch.utils.data import DataLoader\n\n# 将模型设置为评估模式\nmodel.eval()\n# 将模型移动到 GPU 上\nmodel = model.cuda()\n\n# 初始化一个空的列表 submission 用于存储最终的提交结果\nsubmission = []\n\n# 遍历测试数据集中的每一行\nfor _, row in df_test.iterrows():\n    # 复制当前行\n    row = dict(row)\n    \n    # 准备数据 - 切割并加载图像块\n    folder_tiles = extract_image_tiles(\n        os.path.join(DATASET_FOLDER, \"test_images\", f\"{str(row['image_id'])}.png\"),\n        IMAGES_FOLDER, size=2048, scale=0.25)\n    \n    # 创建数据集\n    dataset = TilesFolderDataset(folder_tiles, transforms=VALID_TRANSFORM)\n    \n    # 如果数据集为空，则输出一条消息并将当前行添加到 submission 列表中，继续下一次循环\n    if not len(dataset):\n        print (f\"seem no tiles were cut for `{folder_tiles}`\")\n        submission.append(row)\n        continue\n    \n    # 创建 DataLoader 用于迭代数据集，设置批处理大小为 4，使用 10 个工作进程并关闭 shuffle\n    dataloader = DataLoader(dataset, batch_size=4, num_workers=10, shuffle=False)\n    \n    # 遍历 DataLoader 中的每个批次，用模型进行预测，并将预测结果（概率值）附加到 preds 列表中\n    preds = []\n    for imgs in dataloader:\n        with torch.no_grad():\n            pred = model(imgs.cuda())\n        preds += pred.cpu().numpy().tolist()\n    \n    # 输出当前图像块的总贡献和最大贡献\n    print(f\"Sum contribution from all tiles: {np.sum(preds, axis=0)}\")\n    print(f\"Max contribution over all tiles: {np.max(preds, axis=0)}\")\n    \n    # 决定标签\n    lb = np.argmax(np.sum(preds, axis=0))\n    print(lb)\n    \n    # 更新当前行的标签为对应的标签\n    row['label'] = labels[lb]\n    print(row)\n    \n    # 将更新后的行添加到 submission 列表中\n    submission.append(row)\n    \n    # 清理操作 - 注释掉了删除生成的图像块文件夹的代码，可以选择删除\n    # shutil.rmtree(folder_tiles)\n    os.system(f\"rm -rf {folder_tiles}\")\n\n# 将 submission 转换为 Pandas DataFrame\ndf_sub = pd.DataFrame(submission)\n","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:26:19.830014Z","iopub.execute_input":"2023-11-18T11:26:19.830305Z","iopub.status.idle":"2023-11-18T11:27:16.277895Z","shell.execute_reply.started":"2023-11-18T11:26:19.830281Z","shell.execute_reply":"2023-11-18T11:27:16.276817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(df_sub.head())\ndf_sub[[\"image_id\", \"label\"]].to_csv(\"submission.csv\", index=False)\n\n! head submission.csv","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:27:16.279671Z","iopub.execute_input":"2023-11-18T11:27:16.280604Z","iopub.status.idle":"2023-11-18T11:27:17.274989Z","shell.execute_reply.started":"2023-11-18T11:27:16.280566Z","shell.execute_reply":"2023-11-18T11:27:17.273914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls","metadata":{"pycharm":{"name":"#%%\n"},"execution":{"iopub.status.busy":"2023-11-18T11:27:17.276797Z","iopub.execute_input":"2023-11-18T11:27:17.277771Z","iopub.status.idle":"2023-11-18T11:27:18.261711Z","shell.execute_reply.started":"2023-11-18T11:27:17.277731Z","shell.execute_reply":"2023-11-18T11:27:18.260513Z"},"trusted":true},"execution_count":null,"outputs":[]}]}