{"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":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":1807973,"sourceType":"datasetVersion","datasetId":1074109}],"dockerImageVersionId":30626,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<center><font size=\"6\">SenNet + HOA - Hacking the Human Vasculature in 3D</font></center>\n\n- 赛题链接：https://www.kaggle.com/competitions/blood-vessel-segmentation\n- 赛题类型：计算机视觉、3D语义分割\n- 赛题任务：3D扫描中人体肾脏的分段脉管系统\n\n## 背景介绍\n\n竞赛主办方是Common Fund的细胞衰老网络（SenNet）计划，该计划旨在全面识别和表征体内不同健康状态和寿命阶段的衰老细胞差异。SenNet提供公开可访问的衰老细胞图谱，并开发了建立在单细胞分析的先前进展基础上的创新工具和技术。\n\n目前，人类专家标注员通过手工追踪血管结构，这是一个缓慢的过程。即使有专业标注员，每个新数据集也需要6个月以上的时间才能完成。由于人体解剖结构的变异性以及HiP-CT技术不断改进和变化导致图像质量的变化，使用这些手工数据的机器学习方法在新数据集上泛化效果不佳。\n\n## 赛题任务\n\n在这个竞赛中，你将处理一个数据集，其中包括多个肾脏的高分辨率3D图像，以及它们血管系统的3D分割掩模。你的任务是为测试集中的肾脏数据集创建分割掩模。\n\n## 赛题数据\n\n-   **train/{dataset}/images** - 包含来自多个肾脏数据集的TIFF扫描。每个图像代表3D体积的2D切片，沿着z轴排列，文件从顶部到底部进行枚举。图像切片应该垂直或深度堆叠。\n\n-   **train/{dataset}/labels** - 包含图像的血管分割掩模，格式为TIFF。其中的{dataset}文件夹包括以下内容：\n\n-   -   `kidney_1_dense` - 以50微米分辨率呈现的右肾的整个结构。整个3D动脉血管树已经密集分割，直到距肾小球（即毛细血管床）两代。使用BM05光束线。\n    -   `kidney_1_voi` - `kidney_1`的高分辨率子集，分辨率为5.2微米。\n    -   `kidney_2` - 另一位捐赠者的整个肾脏，以50微米分辨率呈现。稀疏分割（约65%）。\n    -   `kidney_3_dense` - 以50.16微米分辨率捕获的肾脏部分（500张切片），使用BM05光束线。密集分割。请注意，我们在`kidney_3_sparse/images`文件夹中提供了`kidney_3`的所有图像。因此，该数据集仅有一个`labels`文件夹。\n    -   `kidney_3_sparse` - `kidney_3`的其余分割掩模。稀疏分割（约85%）。\n\n-   **test/{dataset}/images** - 包含测试集的TIFF扫描。这些扫描可能与训练集中使用的扫描采用不同的光束线或分辨率。数据集的名称为`kidney_5`和`kidney_6`。\n\n-   **train_rles.csv** - 训练集图像的运行长度编码分割掩模。\n\n-   -   `id` - 每个切片的唯一标识符，格式为`{dataset}_{slice}`。\n    -   `rle` - 该切片的运行长度编码掩模。\n\n-   **sample_submission.csv** - 一个正确格式的示例提交文件。有关详细信息，请参阅评估页面。\n\n## 评价指标\n\n使用容差为 0.0 的表面骰子指标来评估提交的内容。可以在此笔记本中找到该指标的代码：https://www.kaggle.com/metric/surface-dice-metric\n\n```\nid,rle\nkidney_5_0,1 1 100 10\nkidney_5_1,1 1 100 10\nkidney_6_0,1 0\nkidney_6_1,1 0\n```\n\n\n## 赛题赛程\n\n- 2023 年 11 月 7 日 - 开始日期。\n- 2024 年 1 月 30 日 - 报名截止日期。\n- 2024 年 1 月 30 日 - 合并截止日期。\n- 2024 年 2 月 6 日 - 最终提交截止日期。\n\n## 赛题资料\n\n### Shared solutions\n\n- resnet50 2d-unet + xy,zy,zx + cc3d (0.808)\n    - Inference: https://www.kaggle.com/code/hengck23/lb0-808-resnet50-2d-unet-xy-zy-zx-cc3d\n- se_resnext50_32x4d with 2.5d UNet (0.572)\n    - Train: https://www.kaggle.com/code/yoyobar/2-5d-cutting-model-baseline-training/\n    - Inference: https://www.kaggle.com/code/yoyobar/2-5d-cutting-model-baseline-inference/\n- resnext50_32x4d with UNet (0.519)\n    - Inference: https://www.kaggle.com/code/deanphamnguyen/second-sennet-submission-c96555\n- efficientnet-b1 with UNet (0.147)\n    - Train: https://www.kaggle.com/code/kashiwaba/sennet-hoa-train-unet-simple-baseline\n    - Inference: https://www.kaggle.com/code/kashiwaba/sennet-hoa-inference-unet-simple-baseline\n    \n### Useful notebooks\n- [SenNet+HOA | Seg. | PyTorch: Attention-Gated UNet](https://www.kaggle.com/code/aniketkolte04/sennet-hoa-seg-pytorch-attention-gated-unet/)\n- [Fast Surface Dice Computation](https://www.kaggle.com/code/junkoda/fast-surface-dice-computation)\n- [SenNet + HOA - Visualize 3D slice & mask](https://www.kaggle.com/code/sb0702/sennet-hoa-visualize-3d-slice-mask)\n\n### Useful discussion\n- [CV vs public LB](https://www.kaggle.com/competitions/blood-vessel-segmentation/discussion/456714)\n- [[lb0.836] experiment results, hopefully open gold solution till 21-jan-2024](https://www.kaggle.com/competitions/blood-vessel-segmentation/discussion/456118)\n- [[LB 0.726] Experimental Configuration and Results](https://www.kaggle.com/competitions/blood-vessel-segmentation/discussion/456787)\n","metadata":{}},{"cell_type":"markdown","source":"## 数据读取\n\nhttps://www.kaggle.com/code/jirkaborovec/sennet-hoa-rle-decode-encode-demo-submission","metadata":{}},{"cell_type":"code","source":"import os, sys, cv2, glob\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\nfrom skimage import color\nshow_args = dict(vmin=0, interpolation='antialiased', interpolation_stage='rgba') \n\n# 数据集路径\nDATASET_PATH = \"/kaggle/input/blood-vessel-segmentation\"\n\ndef rle_decode(mask_rle: str, img_shape: tuple = None) -> np.ndarray:\n    \"\"\"\n    将 Run-Length Encoding (RLE) 格式的掩模解码为二进制图像数组。\n\n    Parameters:\n        mask_rle (str): RLE 格式的掩模字符串。\n        img_shape (tuple): 图像形状的元组，用于创建相应大小的数组。\n\n    Returns:\n        np.ndarray: 解码后的二进制图像数组。\n    \"\"\"\n    seq = mask_rle.split()\n    starts = np.array(list(map(int, seq[0::2])))\n    lengths = np.array(list(map(int, seq[1::2])))\n    assert len(starts) == len(lengths)\n    ends = starts + lengths\n    img = np.zeros((np.product(img_shape),), dtype=np.uint8)\n    for begin, end in zip(starts, ends):\n        img[begin:end] = 1\n    # 重新调整数组形状\n    img.shape = img_shape\n    return img\n\ndef rle_encode(mask, bg = 0) -> dict:\n    \"\"\"\n    将二进制图像数组编码为 Run-Length Encoding (RLE) 格式的字符串。\n\n    Parameters:\n        mask (np.ndarray): 二进制图像数组。\n        bg (int): 背景像素值，将其排除在编码之外。\n\n    Returns:\n        str: RLE 格式的掩模字符串。\n    \"\"\"\n    vec = mask.flatten()\n    nb = len(vec)\n    where = np.flatnonzero\n    starts = np.r_[0, where(~np.isclose(vec[1:], vec[:-1], equal_nan=True)) + 1]\n    lengths = np.diff(np.r_[starts, nb])\n    values = vec[starts]\n    assert len(starts) == len(lengths) == len(values)\n    rle = []\n    for start, length, val in zip(starts, lengths, values):\n        if val == bg:\n            continue\n        rle += [str(start), length]\n    # 后处理，将编码结果连接为字符串\n    return \" \".join(map(str, rle))","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:37:43.539362Z","iopub.execute_input":"2023-12-15T06:37:43.539632Z","iopub.status.idle":"2023-12-15T06:37:44.124110Z","shell.execute_reply.started":"2023-12-15T06:37:43.539607Z","shell.execute_reply":"2023-12-15T06:37:44.123344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 读取提交样例\ndf_submit = pd.read_csv(os.path.join(DATASET_PATH, \"sample_submission.csv\"))\ndf_submit.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:37:44.125639Z","iopub.execute_input":"2023-12-15T06:37:44.126029Z","iopub.status.idle":"2023-12-15T06:37:44.156996Z","shell.execute_reply.started":"2023-12-15T06:37:44.125984Z","shell.execute_reply":"2023-12-15T06:37:44.156124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 读取训练集标签\ndf_train = pd.read_csv(os.path.join(DATASET_PATH, \"train_rles.csv\"))\ndf_train[[\"dataset\", \"slice\"]] = df_train['id'].str.rsplit(pat='_', n=1, expand=True)\ndf_train.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:37:44.158095Z","iopub.execute_input":"2023-12-15T06:37:44.158360Z","iopub.status.idle":"2023-12-15T06:37:45.375342Z","shell.execute_reply.started":"2023-12-15T06:37:44.158336Z","shell.execute_reply":"2023-12-15T06:37:45.374518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 统计dataset的分布\ndf_train['dataset'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:37:45.377762Z","iopub.execute_input":"2023-12-15T06:37:45.378080Z","iopub.status.idle":"2023-12-15T06:37:45.391624Z","shell.execute_reply.started":"2023-12-15T06:37:45.378051Z","shell.execute_reply":"2023-12-15T06:37:45.390831Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 对每个dataset的样例进行可视化\nfor dataset_name in df_train['dataset'].value_counts().index:\n    for _, spl in df_train[df_train['dataset'] == dataset_name].sample(1).iterrows():\n        p_img = os.path.join(DATASET_PATH, \"train\", spl[\"dataset\"], \"images\", f'{spl[\"slice\"]}.tif')\n        if not os.path.isfile(p_img):\n            continue\n        \n        fig, axarr = plt.subplots(ncols=3, figsize=(12, 6))\n        img = plt.imread(p_img)\n        rle_mask = rle_decode(spl[\"rle\"], img_shape=img.shape)\n        axarr[0].imshow(img, cmap=\"gray\")\n        axarr[0].set_title('Image')\n        axarr[1].imshow(color.label2rgb(rle_mask, img, bg_label=0, bg_color=(1.,1.,1.), alpha=0.25))\n        axarr[1].set_title('Image with color mask')\n        axarr[2].imshow(rle_mask, **show_args)\n        axarr[2].set_title('Mask')\n\n        for i in range(3):\n            axarr[i].set_axis_off()","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:37:48.017463Z","iopub.execute_input":"2023-12-15T06:37:48.017750Z","iopub.status.idle":"2023-12-15T06:37:56.497444Z","shell.execute_reply.started":"2023-12-15T06:37:48.017715Z","shell.execute_reply":"2023-12-15T06:37:56.496570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 基础建模\n\nhttps://www.kaggle.com/code/yoyobar/2-5d-cutting-model-baseline-training/\n\nhttps://www.kaggle.com/code/yoyobar/2-5d-cutting-model-baseline-inference/","metadata":{}},{"cell_type":"code","source":"!pip install segmentation_models_pytorch > /dev/null","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:37:56.498742Z","iopub.execute_input":"2023-12-15T06:37:56.499055Z","iopub.status.idle":"2023-12-15T06:38:15.425888Z","shell.execute_reply.started":"2023-12-15T06:37:56.499004Z","shell.execute_reply":"2023-12-15T06:38:15.424816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn  \nfrom torch.cuda.amp import autocast\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.parallel import DataParallel\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:38:15.427454Z","iopub.execute_input":"2023-12-15T06:38:15.427754Z","iopub.status.idle":"2023-12-15T06:38:22.928190Z","shell.execute_reply.started":"2023-12-15T06:38:15.427725Z","shell.execute_reply":"2023-12-15T06:38:22.927385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    # ============== pred target =============\n    target_size = 1\n\n    # ============== model CFG =============\n    model_name = 'Unet'\n    backbone = 'se_resnext50_32x4d'\n\n    in_chans = 5 # 输入通道数\n    # ============== training CFG =============\n    image_size = 256\n    input_size=256\n    drop_egde_pixel = 0\n    tile_size = image_size\n    stride = tile_size // 2\n    assert stride > drop_egde_pixel  # 确保步幅大于丢弃的边缘像素数\n\n    train_batch_size = 16  # 训练批次大小\n    valid_batch_size = train_batch_size * 2  # 验证批次大小\n\n    epochs = 10  # 训练轮次\n    lr = 5e-4  # 学习率\n\n    # ============== fold =============\n    valid_id = 1  # 验证集标识\n\n    # ============== augmentation =============\n    train_aug_list = [\n        A.RandomResizedCrop(\n            input_size, input_size, scale=(0.8, 1.25)),\n        A.ShiftScaleRotate(p=0.75),\n        A.OneOf([\n                A.GaussNoise(var_limit=[10, 50]),\n                A.GaussianBlur(),\n                A.MotionBlur(),\n                ], p=0.4),\n        A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.5),\n        ToTensorV2(transpose_mask=True),\n    ]\n    train_aug = A.Compose(train_aug_list)  # 训练集数据增强\n\n    valid_aug_list = [\n        ToTensorV2(transpose_mask=True),\n    ]\n    valid_aug = A.Compose(valid_aug_list)  # 验证集数据增强","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:38:22.929396Z","iopub.execute_input":"2023-12-15T06:38:22.929708Z","iopub.status.idle":"2023-12-15T06:38:22.938631Z","shell.execute_reply.started":"2023-12-15T06:38:22.929682Z","shell.execute_reply":"2023-12-15T06:38:22.937634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    def __init__(self, CFG, weight=None):\n        super().__init__()\n        self.CFG = CFG\n        self.encoder = smp.Unet(\n            encoder_name=CFG.backbone, \n            encoder_weights=weight,\n            in_channels=CFG.in_chans,\n            classes=CFG.target_size,\n            activation=None,\n        )\n\n    def forward(self, image):\n        output = self.encoder(image)\n        return output[:,0]","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:44:50.029842Z","iopub.execute_input":"2023-12-15T06:44:50.030608Z","iopub.status.idle":"2023-12-15T06:44:50.037160Z","shell.execute_reply.started":"2023-12-15T06:44:50.030568Z","shell.execute_reply":"2023-12-15T06:44:50.036187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def min_max_normalization(x: torch.Tensor) -> torch.Tensor:\n    \"\"\"\n    对输入张量进行最小-最大归一化处理。\n\n    Parameters:\n        x (torch.Tensor): 输入张量，形状为 (batch, f1, ...)\n\n    Returns:\n        torch.Tensor: 归一化后的张量，形状保持不变。\n    \"\"\"\n    shape = x.shape\n    if x.ndim > 2:\n        x = x.reshape(x.shape[0], -1)\n\n    min_ = x.min(dim=-1, keepdim=True)[0]\n    max_ = x.max(dim=-1, keepdim=True)[0]\n\n    # 如果均值为0，方差为1，表示已经是归一化状态，直接返回\n    if min_.mean() == 0 and max_.mean() == 1:\n        return x.reshape(shape)\n\n    # 进行最小-最大归一化\n    x = (x - min_) / (max_ - min_ + 1e-9)\n    return x.reshape(shape)\n\nclass Data_loader(Dataset):\n    def __init__(self, path, s=\"/images/\"):\n        \"\"\"\n        数据加载器类，用于加载图像数据集。\n\n        Parameters:\n            path (str): 数据集路径。\n            s (str): 图像文件夹名称，默认为\"/images/\"。\n        \"\"\"\n        self.paths = glob.glob(path + f\"{s}*.tif\")\n        self.paths.sort()\n        self.bool = s == \"/labels/\"\n\n    def __len__(self):\n        \"\"\"\n        返回数据集中图像的数量。\n\n        Returns:\n            int: 数据集中图像的数量。\n        \"\"\"\n        return len(self.paths)\n\n    def __getitem__(self, index):\n        \"\"\"\n        获取指定索引处的图像。\n\n        Parameters:\n            index (int): 图像的索引。\n\n        Returns:\n            torch.Tensor: 加载的图像张量。\n        \"\"\"\n        img = cv2.imread(self.paths[index], cv2.IMREAD_GRAYSCALE)\n        img = torch.from_numpy(img)\n\n        # 根据标志将图像转换为相应的数据类型\n        if self.bool:\n            img = img.to(torch.bool)\n        else:\n            img = img.to(torch.uint8)\n\n        return img\n\ndef load_data(path,s):\n    data_loader=Data_loader(path,s)\n    data_loader=DataLoader(data_loader, batch_size=16, num_workers=2)\n    data=[]\n    for x in tqdm(data_loader):\n        data.append(x)\n    return torch.cat(data,dim=0)\n\n#https://www.kaggle.com/code/kashiwaba/sennet-hoa-train-unet-simple-baseline\ndef dice_coef(pred:torch.Tensor,target:torch.Tensor,TH=0.5,epsilon=1e-5):\n    if torch.any(pred<0) or torch.any(pred>1):\n        pred=pred.sigmoid()\n    target = target.unsqueeze(1).to(torch.float32)\n    pred = (pred>TH).to(torch.float32)\n    inter = (target*pred).sum()\n    den = target.sum() + pred.sum()\n    dice = ((2*inter+epsilon)/(den+epsilon)).mean()\n    return dice\n    \nclass Kaggld_Dataset(Dataset):\n    def __init__(self,x:list,y:list,arg=False):\n        super(Dataset,self).__init__()\n        self.x=x#list[(C,H,W),...]\n        self.y=y#list[(C,H,W),...]\n        self.image_size=CFG.image_size\n        self.in_chans=CFG.in_chans\n        self.arg=arg\n        if arg:\n            self.transform=CFG.train_aug\n        else: \n            self.transform=CFG.valid_aug\n\n    def __len__(self) -> int:\n        return sum([y.shape[0]-self.in_chans for y in self.y])\n    \n    def __getitem__(self,index):\n        i=0\n        for x in self.x:\n            if index>x.shape[0]-self.in_chans:\n                index-=x.shape[0]-self.in_chans\n                i+=1\n            else:\n                break\n        x=self.x[i]\n        y=self.y[i]\n        \n        x_index=np.random.randint(0,x.shape[1]-self.image_size)\n        y_index=np.random.randint(0,x.shape[2]-self.image_size)\n\n        x=x[index:index+self.in_chans,x_index:x_index+self.image_size,y_index:y_index+self.image_size].to(torch.float32)\n        y=y[index+self.in_chans//2,x_index:x_index+self.image_size,y_index:y_index+self.image_size].to(torch.float32)\n\n        data = self.transform(image=x.numpy().transpose(1,2,0), mask=y.numpy())\n        x = data['image']\n        y = data['mask']\n        if self.arg:\n            i=np.random.randint(4)\n            x=x.rot90(i,dims=(1,2))\n            y=y.rot90(i,dims=(0,1))\n            for i in range(3):\n                if np.random.randint(2):\n                    x=x.flip(dims=(i,))\n                    if i>=1:\n                        y=y.flip(dims=(i-1,))\n        return x,y","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:38:22.956222Z","iopub.execute_input":"2023-12-15T06:38:22.956549Z","iopub.status.idle":"2023-12-15T06:38:22.979820Z","shell.execute_reply.started":"2023-12-15T06:38:22.956516Z","shell.execute_reply":"2023-12-15T06:38:22.979056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_x = []\ntrain_y = []\n\npaths=glob.glob(os.path.join(DATASET_PATH, \"train/*\"))\npaths.sort()\n\nfor i,path in enumerate(paths[2:]):\n    if \"kidney_3_dense\" in path:\n        continue\n    \n    x = load_data(path,\"/images/\")\n    y = load_data(path,\"/labels/\")\n    train_x.append(x)\n    train_y.append(y)\n\n    #aug\n    train_x.append(x.permute(1,2,0))\n    train_y.append(y.permute(1,2,0))\n    train_x.append(x.permute(2,0,1))\n    train_y.append(y.permute(2,0,1))\n\nval_x=load_data(paths[0],\"/images/\")\nval_y=load_data(paths[0],\"/labels/\")","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:38:22.980942Z","iopub.execute_input":"2023-12-15T06:38:22.981273Z","iopub.status.idle":"2023-12-15T06:42:04.607320Z","shell.execute_reply.started":"2023-12-15T06:38:22.981243Z","shell.execute_reply":"2023-12-15T06:42:04.605256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_x), len(train_y), val_x.shape, val_y.shape","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:42:04.608852Z","iopub.execute_input":"2023-12-15T06:42:04.609371Z","iopub.status.idle":"2023-12-15T06:42:04.616427Z","shell.execute_reply.started":"2023-12-15T06:42:04.609338Z","shell.execute_reply":"2023-12-15T06:42:04.615690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints/\n!cp /kaggle/input/se-net-pretrained-imagenet-weights/* /root/.cache/torch/hub/checkpoints/","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:42:04.617746Z","iopub.execute_input":"2023-12-15T06:42:04.618060Z","iopub.status.idle":"2023-12-15T06:42:19.743051Z","shell.execute_reply.started":"2023-12-15T06:42:04.618028Z","shell.execute_reply":"2023-12-15T06:42:19.741764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CustomModel(CFG, \"imagenet\")\nmodel = model.cuda()","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:42:19.744562Z","iopub.execute_input":"2023-12-15T06:42:19.744841Z","iopub.status.idle":"2023-12-15T06:42:23.634326Z","shell.execute_reply.started":"2023-12-15T06:42:19.744815Z","shell.execute_reply":"2023-12-15T06:42:23.633357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset=Kaggld_Dataset(train_x,train_y,arg=True)\ntrain_dataset = DataLoader(train_dataset, batch_size=16, num_workers=2, shuffle=True, pin_memory=True)\nval_dataset = Kaggld_Dataset([val_x],[val_y])\nval_dataset = DataLoader(val_dataset, batch_size=16, num_workers=2, shuffle=False, pin_memory=True)\n\nloss_fn=nn.BCEWithLogitsLoss()\noptimizer=torch.optim.AdamW(model.parameters(),lr=CFG.lr)\nscaler=torch.cuda.amp.GradScaler()\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer, \n    max_lr=CFG.lr,\n    steps_per_epoch=len(train_dataset), \n    epochs=CFG.epochs+1,\n    pct_start=0.1\n)\n\nfor epoch in range(CFG.epochs):\n    time=tqdm(range(len(train_dataset)))\n    losss=0\n    scores=0\n    for i,(x,y) in enumerate(train_dataset):\n        x=x.cuda()\n        y=y.cuda()\n        x=min_max_normalization(x)\n\n        with autocast():\n            pred=model(x)\n            loss=loss_fn(pred,y)\n\n            optimizer.zero_grad()\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            \n        scheduler.step()\n        score=dice_coef(pred.detach(),y)\n        losss=(losss*i+loss.item())/(i+1)\n        scores=(scores*i+score)/(i+1)\n        time.set_description(f\"epoch:{epoch},loss:{losss:.4f},score:{scores:.4f},lr{optimizer.param_groups[0]['lr']:.4e}\")\n        time.update()\n        \n        del loss,pred\n    \n    time.close()\n    val_losss=0\n    val_scores=0\n    time=tqdm(range(len(val_dataset)))\n    \n    for i,(x,y) in enumerate(val_dataset):\n        x=x.cuda()\n        y=y.cuda()\n        \n        x=min_max_normalization(x)\n\n        with autocast():\n            with torch.no_grad():\n                pred=model(x)\n                loss=loss_fn(pred,y)\n        score=dice_coef(pred.detach(),y)\n        val_losss=(val_losss*i+loss.item())/(i+1)\n        val_scores=(val_scores*i+score)/(i+1)\n        time.set_description(f\"val-->loss:{val_losss:.4f},score:{val_scores:.4f}\")\n        time.update()\n\n    time.close()\n    torch.save(model.state_dict(),f\"./{CFG.backbone}_{epoch}_loss{losss:.2f}_score{scores:.2f}_val_loss{val_losss:.2f}_val_score{val_scores:.2f}.pt\")\n\ntime.close()","metadata":{"execution":{"iopub.status.busy":"2023-12-15T06:42:23.635643Z","iopub.execute_input":"2023-12-15T06:42:23.635938Z","iopub.status.idle":"2023-12-15T06:44:44.482472Z","shell.execute_reply.started":"2023-12-15T06:42:23.635913Z","shell.execute_reply":"2023-12-15T06:44:44.481194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}