{"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":30201,"databundleVersionId":2750748,"sourceType":"competition"},{"sourceId":7269956,"sourceType":"datasetVersion","datasetId":4214254},{"sourceId":7270128,"sourceType":"datasetVersion","datasetId":4214367},{"sourceId":7272905,"sourceType":"datasetVersion","datasetId":4216294}],"dockerImageVersionId":30626,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/pakage/pakage/')","metadata":{"execution":{"iopub.status.busy":"2023-12-25T11:51:20.300651Z","iopub.execute_input":"2023-12-25T11:51:20.301013Z","iopub.status.idle":"2023-12-25T11:51:20.306314Z","shell.execute_reply.started":"2023-12-25T11:51:20.300987Z","shell.execute_reply":"2023-12-25T11:51:20.305270Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import ml_collections","metadata":{"execution":{"iopub.status.busy":"2023-12-25T11:51:20.307831Z","iopub.execute_input":"2023-12-25T11:51:20.308221Z","iopub.status.idle":"2023-12-25T11:51:20.333279Z","shell.execute_reply.started":"2023-12-25T11:51:20.308190Z","shell.execute_reply":"2023-12-25T11:51:20.332157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_b16_config(rate):\n    \"\"\"Returns the ViT-B/16 configuration.\"\"\"\n    # 创建了一个空的配置字典，并将其赋值给变量config\n    config = ml_collections.ConfigDict()\n    config.patches = ml_collections.ConfigDict({'size': (16, 16)})\n    config.hidden_size = 768\n    config.transformer = ml_collections.ConfigDict()\n    config.transformer.mlp_dim = 3072\n    config.transformer.num_heads = 12\n    config.transformer.num_layers = 12\n    config.transformer.attention_dropout_rate = rate\n    config.transformer.dropout_rate = rate\n\n    config.classifier = 'seg'\n    config.representation_size = None\n    config.resnet_pretrained_path = None\n    config.pretrained_path = '../model/vit_checkpoint/imagenet21k/ViT-B_16.npz'\n    config.patch_size = 16\n\n    config.decoder_channels = (256, 128, 64, 16)\n    config.n_classes = 1\n    config.activation = 'softmax'\n    return config\n\n\ndef get_r50_b16_config(rate):\n    \"\"\"Returns the Resnet50 + ViT-B/16 configuration.\"\"\"\n    config = get_b16_config(rate)\n    config.patches.grid = (16, 16)\n    config.resnet = ml_collections.ConfigDict()\n    config.resnet.num_layers = (3, 4, 9)\n    config.resnet.width_factor = 1\n\n    config.classifier = 'seg'\n    config.pretrained_path = '../model/vit_checkpoint/imagenet21k/R50+ViT-B_16.npz'\n    config.decoder_channels = (256, 128, 64, 16)\n    config.skip_channels = [512, 256, 64, 16]\n    config.n_classes = 1\n    config.n_skip = 3\n    config.activation = 'softmax'\n\n    return config","metadata":{"execution":{"iopub.status.busy":"2023-12-25T11:51:20.334312Z","iopub.execute_input":"2023-12-25T11:51:20.334565Z","iopub.status.idle":"2023-12-25T11:51:20.345506Z","shell.execute_reply.started":"2023-12-25T11:51:20.334542Z","shell.execute_reply":"2023-12-25T11:51:20.344658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\n\nfrom os.path import join as pjoin\nfrom collections import OrderedDict\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\n# 将NumPy数组转换为PyTorch张量\ndef np2th(weights, conv=False):\n    \"\"\"Possibly convert HWIO to OIHW.\"\"\"\n    if conv:  # 指示是否执行卷积权重的转置\n        weights = weights.transpose([3, 2, 0, 1])\n    return torch.from_numpy(weights)\n\n\nclass StdConv2d(nn.Conv2d):\n\n    def forward(self, x):\n        w = self.weight  # 类中获取卷积层的权重\n        v, m = torch.var_mean(w, dim=[1, 2, 3], keepdim=True, unbiased=False)\n        w = (w - m) / torch.sqrt(v + 1e-5)\n        return F.conv2d(x, w, self.bias, self.stride, self.padding,\n                        self.dilation, self.groups)\n\n\ndef conv3x3(cin, cout, stride=1, groups=1, bias=False):\n    return StdConv2d(cin, cout, kernel_size=3, stride=stride,\n                     padding=1, bias=bias, groups=groups)\n\n\ndef conv1x1(cin, cout, stride=1, bias=False):\n    return StdConv2d(cin, cout, kernel_size=1, stride=stride,\n                     padding=0, bias=bias)\n\n# 定义残差模块\nclass PreActBottleneck(nn.Module):\n    \"\"\"Pre-activation (v2) bottleneck block.\n    \"\"\"\n\n    def __init__(self, cin, cout=None, cmid=None, stride=1):\n        super().__init__()\n        cout = cout or cin\n        cmid = cmid or cout//4\n\n        # 定义残差网络\n        # 32 就是指定的分组数量，cmid 是输入通道数\n        self.gn1 = nn.GroupNorm(32, cmid, eps=1e-6)\n        self.conv1 = conv1x1(cin, cmid, bias=False)\n        self.gn2 = nn.GroupNorm(32, cmid, eps=1e-6)\n        self.conv2 = conv3x3(cmid, cmid, stride, bias=False)  # Original code has it on conv1!!\n        self.gn3 = nn.GroupNorm(32, cout, eps=1e-6)\n        self.conv3 = conv1x1(cmid, cout, bias=False)\n        self.relu = nn.ReLU(inplace=True)\n\n        # shortcut\n        if (stride != 1 or cin != cout):\n            # Projection also with pre-activation according to paper.\n            self.downsample = conv1x1(cin, cout, stride, bias=False)\n            self.gn_proj = nn.GroupNorm(cout, cout)\n\n    def forward(self, x):\n\n        # Residual branch\n        residual = x\n        if hasattr(self, 'downsample'):\n            residual = self.downsample(x)\n            residual = self.gn_proj(residual)\n\n        # Unit's branch\n        y = self.relu(self.gn1(self.conv1(x)))\n        y = self.relu(self.gn2(self.conv2(y)))\n        y = self.gn3(self.conv3(y))\n\n        y = self.relu(residual + y)\n        return y\n\n    def load_from(self, weights, n_block, n_unit):\n        conv1_weight = np2th(weights[pjoin(n_block, n_unit, \"conv1/kernel\")], conv=True)\n        conv2_weight = np2th(weights[pjoin(n_block, n_unit, \"conv2/kernel\")], conv=True)\n        conv3_weight = np2th(weights[pjoin(n_block, n_unit, \"conv3/kernel\")], conv=True)\n\n        gn1_weight = np2th(weights[pjoin(n_block, n_unit, \"gn1/scale\")])\n        gn1_bias = np2th(weights[pjoin(n_block, n_unit, \"gn1/bias\")])\n\n        gn2_weight = np2th(weights[pjoin(n_block, n_unit, \"gn2/scale\")])\n        gn2_bias = np2th(weights[pjoin(n_block, n_unit, \"gn2/bias\")])\n\n        gn3_weight = np2th(weights[pjoin(n_block, n_unit, \"gn3/scale\")])\n        gn3_bias = np2th(weights[pjoin(n_block, n_unit, \"gn3/bias\")])\n\n        self.conv1.weight.copy_(conv1_weight)\n        self.conv2.weight.copy_(conv2_weight)\n        self.conv3.weight.copy_(conv3_weight)\n\n        self.gn1.weight.copy_(gn1_weight.view(-1))\n        self.gn1.bias.copy_(gn1_bias.view(-1))\n\n        self.gn2.weight.copy_(gn2_weight.view(-1))\n        self.gn2.bias.copy_(gn2_bias.view(-1))\n\n        self.gn3.weight.copy_(gn3_weight.view(-1))\n        self.gn3.bias.copy_(gn3_bias.view(-1))\n\n        if hasattr(self, 'downsample'):\n            proj_conv_weight = np2th(weights[pjoin(n_block, n_unit, \"conv_proj/kernel\")], conv=True)\n            proj_gn_weight = np2th(weights[pjoin(n_block, n_unit, \"gn_proj/scale\")])\n            proj_gn_bias = np2th(weights[pjoin(n_block, n_unit, \"gn_proj/bias\")])\n\n            self.downsample.weight.copy_(proj_conv_weight)\n            self.gn_proj.weight.copy_(proj_gn_weight.view(-1))\n            self.gn_proj.bias.copy_(proj_gn_bias.view(-1))\n\nclass ResNetV2(nn.Module):\n    \"\"\"Implementation of Pre-activation (v2) ResNet mode.\"\"\"\n\n    # block_units 是一个包含三个元素的列表，分别表示每个阶段的残差块的数量。\n    # width_factor 是一个用于控制通道数的缩放因子\n    def __init__(self, block_units=(3, 4, 9), width_factor=1):\n        super().__init__()\n        width = int(64 * width_factor)\n        self.width = width\n\n        self.root = nn.Sequential(OrderedDict([\n            ('conv', StdConv2d(3, width, kernel_size=7, stride=2, bias=False, padding=3)),\n            ('gn', nn.GroupNorm(32, width, eps=1e-6)),\n            ('relu', nn.ReLU(inplace=True)),\n            # ('pool', nn.MaxPool2d(kernel_size=3, stride=2, padding=0))\n        ]))\n\n        self.body = nn.Sequential(OrderedDict([\n            ('block1', nn.Sequential(OrderedDict(\n                [('unit1', PreActBottleneck(cin=width, cout=width*4, cmid=width))] +\n                [(f'unit{i:d}', PreActBottleneck(cin=width*4, cout=width*4, cmid=width)) for i in range(2, block_units[0] + 1)],\n                ))),\n            ('block2', nn.Sequential(OrderedDict(\n                [('unit1', PreActBottleneck(cin=width*4, cout=width*8, cmid=width*2, stride=2))] +\n                [(f'unit{i:d}', PreActBottleneck(cin=width*8, cout=width*8, cmid=width*2)) for i in range(2, block_units[1] + 1)],\n                ))),\n            ('block3', nn.Sequential(OrderedDict(\n                [('unit1', PreActBottleneck(cin=width*8, cout=width*16, cmid=width*4, stride=2))] +\n                [(f'unit{i:d}', PreActBottleneck(cin=width*16, cout=width*16, cmid=width*4)) for i in range(2, block_units[2] + 1)],\n                ))),\n        ]))\n\n    def forward(self, x):\n        features = []  # 用于存储不同阶段的特征\n        b, c, in_size, _ = x.size()  # 获取输入张量的大小信息\n        x = self.root(x)  # 将输入通过初始部分 root 进行处理\n        features.append(x)  # 将处理后的特征添加到列表中\n        x = nn.MaxPool2d(kernel_size=3, stride=2, padding=0)(x)  # 最大池化层进行下采样\n        for i in range(len(self.body)-1):\n            x = self.body[i](x)  # 对 body 中的每个阶段进行前向传播\n            right_size = int(in_size / 4 / (i+1))  # 计算当前阶段的目标尺寸\n            if x.size()[2] != right_size:  # 如果当前特征的尺寸不符合目标尺寸\n                pad = right_size - x.size()[2]  # 计算需要填充的大小\n                assert pad < 3 and pad > 0, \"x {} should {}\".format(x.size(), right_size)\n                feat = torch.zeros((b, x.size()[1], right_size, right_size), device=x.device)\n                feat[:, :, 0:x.size()[2], 0:x.size()[3]] = x[:]\n            else:\n                feat = x\n            features.append(feat)\n        x = self.body[-1](x)  # 对最后一个阶段进行前向传播\n        return x, features[::-1]  # 返回最终输出和各个阶段的特征（反转列表顺序）\n","metadata":{"execution":{"iopub.status.busy":"2023-12-25T11:51:20.347548Z","iopub.execute_input":"2023-12-25T11:51:20.347870Z","iopub.status.idle":"2023-12-25T11:51:20.384428Z","shell.execute_reply.started":"2023-12-25T11:51:20.347841Z","shell.execute_reply":"2023-12-25T11:51:20.383598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# coding=utf-8\nfrom __future__ import absolute_import\nfrom __future__ import division\nfrom __future__ import print_function\n\nimport copy\nimport logging\nimport math\n\nfrom os.path import join as pjoin\n\nimport torch\nimport torch.nn as nn\nimport numpy as np\n\nfrom torch.nn import CrossEntropyLoss, Dropout, Softmax, Linear, Conv2d, LayerNorm\nfrom torch.nn.modules.utils import _pair\nfrom scipy import ndimage\n\n\nlogger = logging.getLogger(__name__)\n\n\nATTENTION_Q = \"MultiHeadDotProductAttention_1/query\"\nATTENTION_K = \"MultiHeadDotProductAttention_1/key\"\nATTENTION_V = \"MultiHeadDotProductAttention_1/value\"\nATTENTION_OUT = \"MultiHeadDotProductAttention_1/out\"\nFC_0 = \"MlpBlock_3/Dense_0\"\nFC_1 = \"MlpBlock_3/Dense_1\"\nATTENTION_NORM = \"LayerNorm_0\"\nMLP_NORM = \"LayerNorm_2\"\n\n\ndef np2th(weights, conv=False):\n    \"\"\"Possibly convert HWIO to OIHW.\"\"\"\n    if conv:\n        weights = weights.transpose([3, 2, 0, 1])\n    return torch.from_numpy(weights)\n\n\n# swish激活函数\ndef swish(x):\n    return x * torch.sigmoid(x)\n\n\nACT2FN = {\"gelu\": torch.nn.functional.gelu, \"relu\": torch.nn.functional.relu, \"swish\": swish}\n\n# 定义注意力机制 输入向量，输出向量\nclass Attention(nn.Module):\n    def __init__(self, config, vis):\n        super(Attention, self).__init__()\n        # 表示是否可视化注意力分布\n        self.vis = vis\n        # 表示注意力头的数量\n        self.num_attention_heads = config.transformer[\"num_heads\"]\n        # 表示每个注意力头的大小，通过总的隐藏层大小除以注意力头的数量计算得到。\n        self.attention_head_size = int(config.hidden_size / self.num_attention_heads)\n        # 表示所有注意力头的总大小\n        self.all_head_size = self.num_attention_heads * self.attention_head_size\n\n        self.query = Linear(config.hidden_size, self.all_head_size)\n        self.key = Linear(config.hidden_size, self.all_head_size)\n        self.value = Linear(config.hidden_size, self.all_head_size)\n        # 用于将注意力计算的结果进行线性变换\n        self.out = Linear(config.hidden_size, config.hidden_size)\n        self.attn_dropout = Dropout(config.transformer[\"attention_dropout_rate\"])\n        self.proj_dropout = Dropout(config.transformer[\"attention_dropout_rate\"])\n\n        self.softmax = Softmax(dim=-1)\n\n    # 调整线性变换的输出，以适应注意力头的形状\n    def transpose_for_scores(self, x):\n        # 计算新的形状，将原来的形状中的最后一个维度拆分成两个维度\n        new_x_shape = x.size()[:-1] + (self.num_attention_heads, self.attention_head_size)\n        # 使用 view 函数调整张量的形状\n        x = x.view(*new_x_shape)\n        # 使用 permute 函数交换维度的顺序\n        return x.permute(0, 2, 1, 3)\n\n    def forward(self, hidden_states):\n        # 通过线性变换得到查询、键、值的表示\n        mixed_query_layer = self.query(hidden_states)\n        mixed_key_layer = self.key(hidden_states)\n        mixed_value_layer = self.value(hidden_states)\n        # 调整查询、键、值的形状以适应注意力头\n        query_layer = self.transpose_for_scores(mixed_query_layer)\n        key_layer = self.transpose_for_scores(mixed_key_layer)\n        value_layer = self.transpose_for_scores(mixed_value_layer)\n        # 计算注意力分数\n        attention_scores = torch.matmul(query_layer, key_layer.transpose(-1, -2))\n        # 缩放注意力分数\n        attention_scores = attention_scores / math.sqrt(self.attention_head_size)\n        # 使用 softmax 函数计算注意力权重\n        attention_probs = self.softmax(attention_scores)\n        # 如果需要可视化，保存注意力权重；否则为 None\n        weights = attention_probs if self.vis else None\n        # 对注意力权重引入 dropout 随机性\n        attention_probs = self.attn_dropout(attention_probs)\n        # 计算加权和的上下文层\n        context_layer = torch.matmul(attention_probs, value_layer)\n        # 调整上下文层的形状\n        context_layer = context_layer.permute(0, 2, 1, 3).contiguous()\n        new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,)\n        context_layer = context_layer.view(*new_context_layer_shape)\n        # 对上下文层进行线性变换\n        attention_output = self.out(context_layer)\n        # 对线性变换的结果引入 dropout 随机性\n        attention_output = self.proj_dropout(attention_output)\n        # 返回最终的注意力输出和注意力权重\n        return attention_output, weights\n\n# 定义多层感知机 输入向量，输出向量\nclass Mlp(nn.Module):\n    def __init__(self, config):\n        super(Mlp, self).__init__()\n        self.fc1 = Linear(config.hidden_size, config.transformer[\"mlp_dim\"])\n        self.fc2 = Linear(config.transformer[\"mlp_dim\"], config.hidden_size)\n        self.act_fn = ACT2FN[\"gelu\"]\n        self.dropout = Dropout(config.transformer[\"dropout_rate\"])\n\n        self._init_weights()\n\n    def _init_weights(self):\n        nn.init.xavier_uniform_(self.fc1.weight)\n        nn.init.xavier_uniform_(self.fc2.weight)\n        nn.init.normal_(self.fc1.bias, std=1e-6)\n        nn.init.normal_(self.fc2.bias, std=1e-6)\n\n    def forward(self, x):\n        x = self.fc1(x)\n        x = self.act_fn(x)\n        x = self.dropout(x)\n        x = self.fc2(x)\n        x = self.dropout(x)\n        return x\n\n# 将特征图进行编码 输入图片，输出编码\nclass Embeddings(nn.Module):\n    \"\"\"Construct the embeddings from patch, position embeddings.\n    \"\"\"\n    def __init__(self, config, img_size, in_channels=3):\n        super(Embeddings, self).__init__()\n        self.hybrid = None  # 是否使用混合嵌入模型\n        self.config = config  # 保存了模型的配置参数\n        img_size = _pair(img_size)  # 将输入的图像尺寸转化为一个二元组，以确保处理不同尺寸的图像\n\n        if config.patches.get(\"grid\") is not None:   # ResNet\n            grid_size = config.patches[\"grid\"]  # 获取配置中的 \"grid\" 参数 （16， 16）\n            patch_size = (img_size[0] // 16 // grid_size[0], img_size[1] // 16 // grid_size[1])\n            patch_size_real = (patch_size[0] * 16, patch_size[1] * 16)\n            n_patches = (img_size[0] // patch_size_real[0]) * (img_size[1] // patch_size_real[1])\n            self.hybrid = True\n        else:\n            patch_size = _pair(config.patches[\"size\"])\n            n_patches = (img_size[0] // patch_size[0]) * (img_size[1] // patch_size[1])\n            self.hybrid = False\n\n        if self.hybrid:\n            self.hybrid_model = ResNetV2(block_units=config.resnet.num_layers, width_factor=config.resnet.width_factor)\n            in_channels = self.hybrid_model.width * 16\n        self.patch_embeddings = Conv2d(in_channels=in_channels,\n                                       out_channels=config.hidden_size,\n                                       kernel_size=patch_size,\n                                       stride=patch_size)\n        self.position_embeddings = nn.Parameter(torch.zeros(1, n_patches, config.hidden_size))\n\n        self.dropout = Dropout(config.transformer[\"dropout_rate\"])\n\n\n    def forward(self, x):\n        if self.hybrid:\n            x, features = self.hybrid_model(x)  # resnet模型\n        else:\n            features = None\n        x = self.patch_embeddings(x)  # (B, hidden. n_patches^(1/2), n_patches^(1/2))\n        x = x.flatten(2)\n        x = x.transpose(-1, -2)  # (B, n_patches, hidden)\n\n        embeddings = x + self.position_embeddings\n        embeddings = self.dropout(embeddings)\n        return embeddings, features\n\n\n# 将编码放入encoder(由Attention和Mlp组成)模块中 输入编码输出特征\nclass Block(nn.Module):\n    def __init__(self, config, vis):\n        super(Block, self).__init__()\n        self.hidden_size = config.hidden_size  # 768\n        self.attention_norm = LayerNorm(config.hidden_size, eps=1e-6)\n        self.ffn_norm = LayerNorm(config.hidden_size, eps=1e-6)\n        self.ffn = Mlp(config)\n        self.attn = Attention(config, vis)\n\n    def forward(self, x):\n        h = x\n        x = self.attention_norm(x)  # 层归一化\n        x, weights = self.attn(x)  # 经过自主意力模块\n        x = x + h  # 残差连接\n\n        h = x\n        x = self.ffn_norm(x)  # 层归一化\n        x = self.ffn(x)  # 经过mlp层\n        x = x + h  # 残差链接\n        return x, weights  # 返回最后的输出和权重\n\n\n\n# 将多层集成块合到一起\nclass Encoder(nn.Module):\n    def __init__(self, config, vis):\n        super(Encoder, self).__init__()\n        self.vis = vis\n        self.layer = nn.ModuleList()\n        self.encoder_norm = LayerNorm(config.hidden_size, eps=1e-6)\n        for _ in range(config.transformer[\"num_layers\"]):  # 循环12层encoder模块\n            layer = Block(config, vis)\n            self.layer.append(copy.deepcopy(layer))\n\n    def forward(self, hidden_states):\n        attn_weights = []\n        for layer_block in self.layer:\n            hidden_states, weights = layer_block(hidden_states)\n            if self.vis:\n                attn_weights.append(weights)\n        encoded = self.encoder_norm(hidden_states)\n        return encoded, attn_weights\n\n# 先进行编码 再放入encoder块中\nclass Transformer(nn.Module):\n    def __init__(self, config, img_size, vis):\n        super(Transformer, self).__init__()\n        self.embeddings = Embeddings(config, img_size=img_size)\n        self.encoder = Encoder(config, vis)\n\n    def forward(self, input_ids):\n        embedding_output, features = self.embeddings(input_ids)\n        encoded, attn_weights = self.encoder(embedding_output)  # (B, n_patch, hidden)\n        return encoded, attn_weights, features\n\n\nclass Conv2dReLU(nn.Sequential):\n    def __init__(\n            self,\n            in_channels,\n            out_channels,\n            kernel_size,\n            padding=0,\n            stride=1,\n            use_batchnorm=True,\n    ):\n        conv = nn.Conv2d(\n            in_channels,\n            out_channels,\n            kernel_size,\n            stride=stride,\n            padding=padding,\n            bias=not (use_batchnorm),\n        )\n        relu = nn.ReLU(inplace=True)\n\n        bn = nn.BatchNorm2d(out_channels)\n\n        super(Conv2dReLU, self).__init__(conv, bn, relu)\n\n# 解码器\nclass DecoderBlock(nn.Module):\n    def __init__(\n            self,\n            in_channels,\n            out_channels,\n            skip_channels=0,\n            use_batchnorm=True,\n    ):\n        super().__init__()\n        self.conv1 = Conv2dReLU(\n            in_channels + skip_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        self.conv2 = Conv2dReLU(\n            out_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        self.up = nn.UpsamplingBilinear2d(scale_factor=2)\n\n    def forward(self, x, skip=None):\n        x = self.up(x)  # 先进行上采样\n        if skip is not None:\n            x = torch.cat([x, skip], dim=1)  # 与之前层进行连接\n        x = self.conv1(x)  # 进行第一层卷积操作\n        x = self.conv2(x)  # 进行第二层卷积操作\n        return x\n\n\nclass SegmentationHead(nn.Sequential):\n\n    def __init__(self, in_channels, out_channels, kernel_size=3, upsampling=1):\n        conv2d = nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, padding=kernel_size // 2)\n        upsampling = nn.UpsamplingBilinear2d(scale_factor=upsampling) if upsampling > 1 else nn.Identity()\n        super().__init__(conv2d, upsampling)\n\n\nclass DecoderCup(nn.Module):\n    def __init__(self, config):\n        super().__init__()\n        self.config = config\n        head_channels = 512\n        self.conv_more = Conv2dReLU(\n            config.hidden_size,\n            head_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=True,\n        )\n        decoder_channels = config.decoder_channels\n        in_channels = [head_channels] + list(decoder_channels[:-1])\n        out_channels = decoder_channels\n\n        if self.config.n_skip != 0:\n            skip_channels = self.config.skip_channels\n            for i in range(4-self.config.n_skip):  # re-select the skip channels according to n_skip\n                skip_channels[3-i]=0\n\n        else:\n            skip_channels=[0,0,0,0]\n\n        blocks = [\n            DecoderBlock(in_ch, out_ch, sk_ch) for in_ch, out_ch, sk_ch in zip(in_channels, out_channels, skip_channels)\n        ]\n        self.blocks = nn.ModuleList(blocks)\n\n    def forward(self, hidden_states, features=None):\n        B, n_patch, hidden = hidden_states.size()  # reshape from (B, n_patch, hidden) to (B, h, w, hidden)\n        h, w = int(np.sqrt(n_patch)), int(np.sqrt(n_patch))\n        x = hidden_states.permute(0, 2, 1)\n        x = x.contiguous().view(B, hidden, h, w)\n        x = self.conv_more(x)\n        for i, decoder_block in enumerate(self.blocks):\n            if features is not None:\n                skip = features[i] if (i < self.config.n_skip) else None\n            else:\n                skip = None\n            x = decoder_block(x, skip=skip)\n        return x\n\n\nclass VisionTransformer(nn.Module):\n    def __init__(self, config, img_size=224, num_classes=21843, zero_head=False, vis=False):\n        super(VisionTransformer, self).__init__()\n        self.num_classes = num_classes\n        self.zero_head = zero_head\n        self.classifier = config.classifier\n        self.transformer = Transformer(config, img_size, vis)\n        self.decoder = DecoderCup(config)\n        self.segmentation_head = SegmentationHead(\n            in_channels=config['decoder_channels'][-1],\n            out_channels=config['n_classes'],\n            kernel_size=3,\n        )\n        self.config = config\n\n    def forward(self, x):\n        if x.size()[1] == 1:\n            x = x.repeat(1,3,1,1)\n        x, attn_weights, features = self.transformer(x)  # (B, n_patch, hidden)\n        x = self.decoder(x, features)\n        logits = self.segmentation_head(x)\n        return logits\n\n\n\n\n\n","metadata":{"execution":{"iopub.status.busy":"2023-12-25T11:51:20.436391Z","iopub.execute_input":"2023-12-25T11:51:20.436691Z","iopub.status.idle":"2023-12-25T11:51:20.494014Z","shell.execute_reply.started":"2023-12-25T11:51:20.436665Z","shell.execute_reply":"2023-12-25T11:51:20.492977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# coding=utf-8\nfrom __future__ import absolute_import\nfrom __future__ import division\nfrom __future__ import print_function\n\nimport math\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport time\nimport random\nimport gc\nfrom pathlib import Path\nimport cv2\nfrom tqdm import tqdm\nfrom sklearn.model_selection import KFold\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n\nimport matplotlib.pyplot as plt\n\nRUN_EDA = True  # EDA数据分析标识\nRUN_TRAINING = False  # 训练标识\nTRAIN_ALL = False  # True: 将所有数据进行训练并输出一个模型. False：交叉验证并输出FOLD_NUM个模型.\nFOLD_NUM = 5  # 交叉验证模型数\nEPOCHS = 20  # epoch\nRUN_INFERENCE = True  # 测试验证标识\n\n# 目录设置\nDATA_DIR = '/kaggle/input/sartorius-cell-instance-segmentation/'\nMODEL_DIR = '/kaggle/input/weight/'\nWEIGHT_DIR = '/kaggle/working/'\nIMG_SAVE_DIR = '/kaggle/working/result/'\n\n# PyTorch变量\nSEED = 42\nNUM_WORKERS = 2\nBATCH_SIZE = 8\nWEIGHT_DECAY = 0.0001\nLR = 0.0001\nMOMENTUM = 0.9\n\n# 阈值\nTHRESHOLD = 0\n\n# 运行设备\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\n\n# 设置种子\nseed_everything(SEED)\n\n\n# 显示GPU使用情况\ndef show_gpu_memory(device):\n    print(f\"Allocated GPU memory: {torch.cuda.memory_allocated(device) / 1024 / 1024:.2f} MB\")\n    print(f\"Cached GPU memory: {torch.cuda.memory_cached(device) / 1024 / 1024:.2f} MB\")\n\n\n# 加载图像\ndef load_img(path):\n    img_bgr = cv2.imread(path)\n    img_rgb = img_bgr[:, :, ::-1]\n    return img_rgb\n\n\n# 将原始dataframe数据打包 并统计统一图像的RLE总数\ndef group_bboxes(df):\n    df_ = df.copy()\n    df_['segment_count'] = 1\n    df_ = df_.groupby(['id', 'width', 'height', 'cell_type']).count().reset_index()\n    return_df = df_[['id', 'width', 'height', 'cell_type', 'segment_count']]\n    return return_df\n\n\n# 将多幅图像重构至Gallery Style\ndef create_gallery(array, ncols=3):\n    \"\"\"\n    Source: https://www.amazon.co.jp/Data-Analysis-Machine-Learning-Kaggle-ebook/dp/B09F3STL34/\n\n    Args:\n        array (numpy.ndarray): array of images.\n        ncols (int, optional): Num of columns. Defaults to 3.\n\n    Returns:\n        numpy.ndarray: One concatenated image.\n    \"\"\"\n    nindex, height, width, intensity = array.shape\n    nrows = nindex // ncols\n    assert nindex == nrows * ncols\n    result = (array.reshape(nrows, ncols, height, width, intensity)\n              .swapaxes(1, 2)\n              .reshape(height * nrows, width * ncols, intensity))\n    return result\n\n\n# 解码RLE\ndef decode_rle(rle, height, width, brightness=1):\n    \"\"\"\n    modified from: https://www.kaggle.com/paulorzp/run-length-encode-and-decode\n\n    Args:\n        rle (str): mask with run length encoding.\n        height (int): return image height.\n        width (int): return image width.\n        brightness (int): brightness of the pixel. Default to 1.\n\n    Returns:\n        np.ndarray: 1(b) - mask, 0 - background.\n    \"\"\"\n    s = rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(height * width, dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = brightness  # 像素亮度\n    return img.reshape((height, width))\n\n\n# 由RLE构造蒙版图像\ndef create_mask_image(image, masks, b=1):\n    \"\"\"\n    Args:\n        image (numpy.ndarray): array of images.\n        masks (list): List with RLE-encoded mask information.\n        b (int): brightness of the pixel. Default to 1.\n\n    Returns:\n        numpy.ndarray: 1(b) - mask, 0 - background.\n    \"\"\"\n\n    s = image.shape\n    h = s[0]\n    w = s[1]\n    mask_image = np.zeros((h, w))\n    for mask in masks:\n        mask_image += decode_rle(mask, h, w, b)\n    mask_image = mask_image.clip(0, 1)\n    return mask_image\n\n\n# 显示验证结果\ndef show_validation_score(train_loss_list, valid_loss_list, valid_iou_list, save=True):\n    save_dir=IMG_SAVE_DIR\n    save_name=f'segmentation_validation_score.png'\n    fig = plt.figure(figsize=(10, 10))\n    for i in range(FOLD_NUM):\n        train_loss = train_loss_list[i]\n        valid_loss = valid_loss_list[i]\n        valid_iou = valid_iou_list[i]\n        ax = fig.add_subplot(math.ceil(np.sqrt(FOLD_NUM)), math.ceil(np.sqrt(FOLD_NUM)), i + 1, title=f'Fold {i + 1}')\n        ax.plot(range(EPOCHS), train_loss, c='orange', label='train')\n        ax.plot(range(EPOCHS), valid_loss, c='blue', label='valid')\n        ax.plot(range(EPOCHS), valid_iou, c='red', label='iou')                     \n        ax.set_xlabel('epoch')\n        ax.set_ylabel('loss')\n        ax.legend()\n    plt.tight_layout()\n    if save:\n        os.makedirs(save_dir, exist_ok=True)\n        plt.savefig(save_dir + save_name)\n        print('img save successfully!')\n    else:\n        plt.show()\n\n\n# 编码为RLE 以空格分隔\ndef encode_rle(predicted_img):\n    predicted_img = (predicted_img > THRESHOLD).astype(int)\n    height, width = predicted_img.shape\n\n    # Get the index of the masked pixel\n    pixels = predicted_img.copy()\n    pixels_list = []\n    for y in range(height):\n        for x in range(width):\n            if pixels[y][x] != 0:\n                pixels_list.append(y * width + x)\n\n    # RLE encoding\n    rle_list = []\n    start = pixels_list[0]\n    count = 1\n    for i in range(1, len(pixels_list)):\n        if pixels_list[i] == pixels_list[i - 1] + 1:\n            count += 1\n        else:\n            rle_list.extend([start, count])\n            start = pixels_list[i]\n            count = 1\n    rle_list.extend([start, count])\n\n    rle_str = [str(x) for x in rle_list]\n    return ' '.join(rle_str)\n\n\n# 训练集图像转换和图像数据增强\ndef transform_train():\n    transforms = [\n        A.Resize(224, 224, p=1),\n        A.HorizontalFlip(p=0.5),\n        A.Transpose(p=0.5),\n        ToTensorV2(p=1)\n    ]\n    return A.Compose(transforms)\n\n\n# 验证集只进行图像转换\ndef transform_valid():\n    transforms = [\n        A.Resize(224, 224, p=1),\n        ToTensorV2(p=1)\n    ]\n    return A.Compose(transforms)\n\n\n# Dataset的子类CellDataset，用于生成DataLoader\nclass CellDataset(Dataset):\n    def __init__(self, image_ids, dataframe, data_root, transforms=None, stage='train'):\n        super().__init__()\n        self.image_ids = image_ids\n        self.dataframe = dataframe\n        self.data_root = data_root\n        self.transforms = transforms\n        self.stage = stage\n\n    def __len__(self):\n        return self.image_ids.shape[0]\n\n    def __getitem__(self, index):\n        image_id = self.image_ids[index]\n        # 加载图像\n        image = load_img(f'{self.data_root}{image_id}.png').astype(np.float32)\n\n        # 3通道转换为1通道，用作训练数据\n        image = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)\n        image /= 255.0  # 归一化\n\n        # 为训练集、验证集和测试集定义不同的数据获取方式\n        # 训练集和验证集\n        if self.stage == 'train':\n            # 由RLE编码生成蒙版图像，用作label\n            masks = self.dataframe[self.dataframe['id'] == image_id]['annotation'].tolist()\n            mask_image = create_mask_image(image, masks)\n\n            # 使用转换关系和图像增强转换图像\n            if self.transforms:\n                transformed = self.transforms(image=image, mask=mask_image)\n                image, mask_image = transformed['image'], transformed['mask']\n            return image, mask_image, image_id\n\n        # 测试集\n        else:\n            # 转换图像\n            if self.transforms:\n                image = self.transforms(image=image)['image']\n\n            return image, image_id\n\n\n# 生成训练集和验证集的DataLoader\ndef create_dataloader(grouped_df, df, trn_idx, val_idx):\n    train_ = grouped_df.loc[trn_idx,:].reset_index(drop=True)\n    valid_ = grouped_df.loc[val_idx,:].reset_index(drop=True)\n\n    # Dataset\n    train_datasets = CellDataset(train_['id'].to_numpy(), df, DATA_DIR+'train/', transforms=transform_train())\n    valid_datasets = CellDataset(valid_['id'].to_numpy(), df, DATA_DIR+'train/', transforms=transform_valid())\n\n    # DataLoader\n    train_loader = DataLoader(train_datasets, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS, shuffle=True)\n    valid_loader = DataLoader(valid_datasets, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS, shuffle=False)\n\n    return train_loader, valid_loader\n\n\ndef IOU(mask, output):\n    # 将掩码转换为二进制形式\n    threshold = THRESHOLD\n    binary_mask1 = (mask > threshold).type(torch.int)\n    binary_mask2 = (output > threshold).type(torch.int)\n\n    # 计算交集和并集\n    intersection = torch.sum(binary_mask1 * binary_mask2, dim=(1, 2, 3))\n    union = torch.sum(binary_mask1 + binary_mask2 - (binary_mask1 * binary_mask2), dim=(1, 2, 3))\n\n    # 计算交并比\n    iou = intersection / union.where(union > 0, torch.tensor(1.0))  # 避免除以零，将零除以非零值得到零\n\n    return iou.mean().item()  # 计算整个批次的平均IoU，并将结果转换为标量值","metadata":{"execution":{"iopub.status.busy":"2023-12-25T12:29:21.378036Z","iopub.execute_input":"2023-12-25T12:29:21.378785Z","iopub.status.idle":"2023-12-25T12:29:21.429790Z","shell.execute_reply.started":"2023-12-25T12:29:21.378741Z","shell.execute_reply":"2023-12-25T12:29:21.428675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == '__main__':\n    df = pd.read_csv(DATA_DIR + 'train.csv')\n    print(df.head())\n\n    grouped_df = group_bboxes(df)\n    print(grouped_df.head())\n\n\n    if RUN_EDA:\n        img_names = Path(DATA_DIR + 'train/').glob('*.png')\n        img_list = []\n        for i, img_name in enumerate(img_names):\n            img_list.append(load_img(img_name.as_posix()))\n            print(img_name.name)\n            if i == 5:\n                break\n        # plt.figure(figsize=(10, 10))\n        # plt.imshow(create_gallery(np.array(img_list), ncols=3))\n    else:\n        print('RUN_EDA is False')\n\n    if RUN_EDA:\n        image_id = grouped_df['id'][0]\n        img = load_img(f'{DATA_DIR}train/{image_id}.png')\n        masks = df[df['id'] == image_id]['annotation'].tolist()\n        masked_img = create_mask_image(img, masks)\n        # plt.figure()\n        # plt.imshow(img)\n        # plt.figure()\n        # plt.imshow(masked_img)\n    else:\n        print('RUN_EDA is False')\n\n    if RUN_EDA:\n        img_shape = set()\n        img_ext = set()\n        img_names = Path(DATA_DIR + 'train/').glob('*')\n        pbar = tqdm(img_names, total=len(grouped_df))\n        for img_name in pbar:\n            img = load_img(img_name.as_posix())\n            img_shape.add(img.shape)\n            img_ext.add(img_name.suffix)\n        print(f'Image shapes are {img_shape}.')\n        print(f'Image extensions are {img_ext}.')\n\n    if RUN_EDA:\n        img_names = Path(DATA_DIR + 'train/').glob('*')\n        # plt.figure(figsize=(10, 10))\n        pbar = tqdm(img_names, total=len(grouped_df))\n        for img_name in pbar:\n            img = load_img(img_name.as_posix())\n            hist = cv2.calcHist([img], [0], None, [256], [0, 256])\n            # plt.plot(hist)\n        # plt.show()\n    else:\n        print('RUN_EDA is False')","metadata":{"execution":{"iopub.status.busy":"2023-12-25T12:32:50.059724Z","iopub.execute_input":"2023-12-25T12:32:50.060723Z","iopub.status.idle":"2023-12-25T12:33:00.147276Z","shell.execute_reply.started":"2023-12-25T12:32:50.060683Z","shell.execute_reply":"2023-12-25T12:33:00.146020Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"    train_loss_list_total = []\n    valid_loss_list_total = []\n    valid_iou_list_total = []\n \n    if RUN_TRAINING:\n        if TRAIN_ALL:\n            # 将所有数据用作训练，并输出一个模型\n            folds = [['', '']]\n        else:\n            # 进行交叉验证\n            folds = KFold(n_splits=FOLD_NUM, shuffle=True, random_state=SEED) \\\n                .split(np.arange(grouped_df.shape[0]), grouped_df['id'].to_numpy())\n\n        # 可视化用\n        train_loss_list = []\n        valid_loss_list = []\n        valid_iou_list = []\n\n        for fold, (trn_idx, val_idx) in enumerate(folds):\n            # 加载数据\n            if TRAIN_ALL:\n                train_datasets = CellDataset(grouped_df['id'].to_numpy(), df, DATA_DIR + 'train/',\n                                             transforms=transform_train())\n                train_loader = DataLoader(train_datasets, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS, shuffle=True)\n            else:\n                print(f'==========Cross-Validation Fold {fold + 1}==========')\n                train_loader, valid_loader = create_dataloader(grouped_df, df, trn_idx, val_idx)\n                # 可视化用\n                valid_losses = []\n                valid_ious = []\n\n            train_losses = []\n\n            # 加载模型、损失函数和优化器\n            # model = UNet().to(device)\n\n            config = get_r50_b16_config(0.5)\n            config.patches.grid = (int(224 / 16), int(224 / 16))\n            model = VisionTransformer(config, img_size=224).to(device)\n            model.load_state_dict(torch.load(MODEL_DIR + 'segmentation(0.5_50).pth'))\n            criterion = nn.BCEWithLogitsLoss().to(device)\n            optimizer = optim.Adam(model.parameters())\n\n            # 训练\n            best_loss = 10 ** 5\n            for epoch in range(EPOCHS):\n                time_start = time.time()\n                print(f'==========Epoch {epoch + 1} Start Training==========')\n                model.train()\n                train_loss = 0\n                pbar = tqdm(enumerate(train_loader), total=len(train_loader))\n                for step, (imgs, masks, image_ids) in pbar:\n                    imgs = imgs.to(device).float()\n                    # imgs = torch.squeeze(imgs)\n                    masks = masks.to(device).float()\n                    masks = masks.view(imgs.shape[0], -1, 224, 224)\n\n                    optimizer.zero_grad()\n                    output = model(imgs)\n\n                    loss = criterion(output, masks)\n                    loss.backward()\n                    optimizer.step()\n                    # print(loss.item())\n                    train_loss += loss.item()\n\n                train_loss /= len(train_loader)\n\n                # 验证\n                if TRAIN_ALL == False:\n                    print(f'==========Epoch {epoch + 1} Start Validation==========')\n\n                    with torch.no_grad():\n                        valid_loss = 0\n                        valid_iou = 0\n                        preds = []\n                        pbar = tqdm(enumerate(valid_loader), total=len(valid_loader))\n                        for step, (imgs, masks, image_ids) in pbar:\n                            imgs = imgs.to(device).float()\n                            # imgs = torch.squeeze(imgs)\n                            masks = masks.to(device).float()\n                            masks = masks.view(imgs.shape[0], -1, 224, 224)\n\n                            val_output = model(imgs)\n                            val_loss = criterion(val_output, masks)\n                            val_iou = IOU(val_output, masks)\n                            valid_loss += val_loss.item()\n                            valid_iou += val_iou\n                        valid_loss /= len(valid_loader)\n                        valid_iou /= len(valid_loader)\n\n                # 输出结果\n                exec_t = int((time.time() - time_start) / 60)\n                if TRAIN_ALL:\n                    print(f'Epoch : {epoch + 1} - loss : {train_loss:.4f} / Exec time {exec_t} min\\n')\n\n                else:\n                    print(\n                        f'Epoch : {epoch + 1} - loss : {train_loss:.4f} - val_loss : {valid_loss:.4f} - iou : {valid_iou:.4f}/ Exec time {exec_t} min\\n'\n                    )\n                    # 存储每个epoch的loss\n                    train_losses.append(train_loss)\n                    valid_losses.append(valid_loss)\n                    valid_ious.append(valid_iou)\n\n            if TRAIN_ALL:\n                print(f'Save model trained with all data')\n                os.makedirs(MODEL_DIR, exist_ok=True)\n                torch.save(model.state_dict(), WEIGHT_DIR + 'segmentation(10).pth')\n                # del model, optimizer, train_loader\n            else:\n                train_loss_list.append(train_losses)\n                valid_loss_list.append(valid_losses)\n                valid_iou_list.append(valid_ious)\n                # del model, optimizer, train_loader, valid_loader, train_losses, valid_losses\n            gc.collect()\n            torch.cuda.empty_cache()\n            \n            break\n        \n        if TRAIN_ALL == False:\n            show_validation_score(train_loss_list, valid_loss_list, valid_iou_list)\n\n        train_loss_list_total.append(train_loss_list)\n        valid_loss_list_total.append(valid_loss_list)\n        valid_iou_list_total.append(valid_iou_list)\n    else:\n        print('RUN_TRAINING is False')\n        \n        \n\n    if RUN_INFERENCE:\n        files = os.listdir(DATA_DIR + 'test/')\n        image_ids = np.array([os.path.splitext(file)[0] for file in files])\n        ids = []\n        rle_test_preds = []\n        original_size = (704, 520)  # (width, height)\n\n        # 加载数据\n        test_datasets = CellDataset(image_ids, df, DATA_DIR + 'test/', transforms=transform_valid(), stage='test')\n\n        # 生成测试数据集DataLoader\n        test_loader = DataLoader(test_datasets, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS, shuffle=False)\n        # 加载模型\n        config = get_r50_b16_config(0.5)\n        config.patches.grid = (int(224 / 16), int(224 / 16))\n        model = VisionTransformer(config, img_size=224).to(device)\n        model.load_state_dict(torch.load(MODEL_DIR + 'segmentation(0.5_50).pth'))\n\n        # 开始测试\n        print(f'==========Start Inference==========')\n        with torch.no_grad():\n            test_preds = []\n            pbar = tqdm(enumerate(test_loader), total=len(test_loader))\n            for step, (imgs, image_ids) in pbar:\n                imgs = imgs.to(device).float()\n                output = model(imgs)\n    \n                # 将output由tensor转换为np.array\n                output = output.detach().cpu().numpy()\n                # 进行RLE编码\n                for image_id, predicted_mask in zip(image_ids, output):\n                    predicted_mask = np.squeeze(predicted_mask)\n\n                    # resize\n                    predicted_mask = cv2.resize(predicted_mask, original_size)\n\n                    rle_mask = encode_rle(predicted_mask)\n                    ids.append(image_id)\n                    rle_test_preds.append(rle_mask)\n\n        submission_df = pd.DataFrame({\n            'id': ids, 'predicted': rle_test_preds\n        })\n        print(submission_df)\n\n        #plt.imshow(predicted_mask > THRESHOLD)\n\n        submission_df.to_csv('submission.csv', index=False)\n\n        target = submission_df.iloc[0]\n        img = load_img(f'{DATA_DIR}test/{target[\"id\"]}.png')\n\n        masks = [target['predicted']]\n        masked_img = create_mask_image(img, masks)\n        plt.figure()\n        plt.imshow(img)\n        plt.figure()\n        # img_color = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)\n        plt.imshow(masked_img)\n        plt.show()\n        print('yes')\n\n    else:\n        print('RUN_INFERENCE is False')","metadata":{"execution":{"iopub.status.busy":"2023-12-25T12:31:18.855932Z","iopub.execute_input":"2023-12-25T12:31:18.856303Z","iopub.status.idle":"2023-12-25T12:31:22.168671Z","shell.execute_reply.started":"2023-12-25T12:31:18.856274Z","shell.execute_reply":"2023-12-25T12:31:22.167728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), WEIGHT_DIR + 'segmentation(10).pth')","metadata":{"execution":{"iopub.status.busy":"2023-12-25T12:23:56.975339Z","iopub.status.idle":"2023-12-25T12:23:56.975662Z","shell.execute_reply.started":"2023-12-25T12:23:56.975504Z","shell.execute_reply":"2023-12-25T12:23:56.975519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(np.shape(train_loss_list_total))","metadata":{"execution":{"iopub.status.busy":"2023-12-25T11:51:36.440346Z","iopub.status.idle":"2023-12-25T11:51:36.440653Z","shell.execute_reply.started":"2023-12-25T11:51:36.440498Z","shell.execute_reply":"2023-12-25T11:51:36.440513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\"\"\nsave_dir=IMG_SAVE_DIR\nx = [p for p in range(1, EPOCHS + 1)]\nfor i in range(2):\n    plt.plot(x, train_loss_list_total[i][0], label = f'Loss = {losses[i]}')\n    plt.title('Train_Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Train_Loss')\nplt.legend()\nsave_name = f'Train_Loss.png'\nplt.savefig(save_dir + save_name)\nplt.show()\n\nfor i in range(2):   \n    plt.plot(x, valid_loss_list_total[i][0], label = f'Loss = {losses[i]}')\n    plt.title('Valid_Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Valid_Loss')\nplt.legend()\nsave_name = f'Valid_Loss.png'\nplt.savefig(save_dir + save_name)\nplt.show()\n\nfor i in range(2):\n    plt.plot(x, valid_iou_list_total[i][0], label = f'Loss = {losses[i]}')\n    plt.title('Valid_IOU')\nplt.xlabel('Epoch')\nplt.ylabel('Valid_IOU')\nplt.legend()\nsave_name = f'Valid_IOU.png'\nplt.savefig(save_dir + save_name)\nprint('img save successfully!')\nplt.show()\n\"\"\"\"\"","metadata":{"execution":{"iopub.status.busy":"2023-12-25T11:51:36.442379Z","iopub.status.idle":"2023-12-25T11:51:36.442838Z","shell.execute_reply.started":"2023-12-25T11:51:36.442595Z","shell.execute_reply":"2023-12-25T11:51:36.442615Z"},"trusted":true},"execution_count":null,"outputs":[]}]}