{"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":7278237,"sourceType":"datasetVersion","datasetId":4219818},{"sourceId":7278310,"sourceType":"datasetVersion","datasetId":4219865},{"sourceId":7283446,"sourceType":"datasetVersion","datasetId":4223319},{"sourceId":7285396,"sourceType":"datasetVersion","datasetId":4224637}],"dockerImageVersionId":30627,"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-27T10:57:08.114378Z","iopub.execute_input":"2023-12-27T10:57:08.115061Z","iopub.status.idle":"2023-12-27T10:57:08.124637Z","shell.execute_reply.started":"2023-12-27T10:57:08.115026Z","shell.execute_reply":"2023-12-27T10:57:08.123768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import ml_collections","metadata":{"execution":{"iopub.status.busy":"2023-12-27T10:57:08.126542Z","iopub.execute_input":"2023-12-27T10:57:08.126835Z","iopub.status.idle":"2023-12-27T10:57:08.243710Z","shell.execute_reply.started":"2023-12-27T10:57:08.126810Z","shell.execute_reply":"2023-12-27T10:57:08.242602Z"},"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-27T10:57:08.244825Z","iopub.execute_input":"2023-12-27T10:57:08.245171Z","iopub.status.idle":"2023-12-27T10:57:08.255505Z","shell.execute_reply.started":"2023-12-27T10:57:08.245144Z","shell.execute_reply":"2023-12-27T10:57:08.254591Z"},"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]  # 返回最终输出和各个阶段的特征（反转列表顺序）","metadata":{"execution":{"iopub.status.busy":"2023-12-27T10:57:08.256646Z","iopub.execute_input":"2023-12-27T10:57:08.256940Z","iopub.status.idle":"2023-12-27T10:57:10.958478Z","shell.execute_reply.started":"2023-12-27T10:57:08.256910Z","shell.execute_reply":"2023-12-27T10:57:10.957538Z"},"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","metadata":{"execution":{"iopub.status.busy":"2023-12-27T10:57:10.960859Z","iopub.execute_input":"2023-12-27T10:57:10.961236Z","iopub.status.idle":"2023-12-27T10:57:11.187934Z","shell.execute_reply.started":"2023-12-27T10:57:10.961209Z","shell.execute_reply":"2023-12-27T10:57:11.186616Z"},"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\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 = True  # 训练标识\nTRAIN_ALL = True  # True: 将所有数据进行训练并输出一个模型. False：交叉验证并输出FOLD_NUM个模型.\nFOLD_NUM = 1  # 交叉验证模型数\nEPOCHS = 200  # epoch\nRUN_INFERENCE = False  # 测试验证标识\n\n# 目录设置\nDATA_DIR = '/kaggle/input/sartorius-cell-instance-segmentation/'\nMODEL_DIR = '/kaggle/input/0-5-100-weight/'\nWEIGHT_DIR = './'\nIMG_SAVE_DIR = '/kaggle/working/'\n\n# PyTorch变量\nSEED = 42\nNUM_WORKERS = 2\nBATCH_SIZE = 8\nWEIGHT_DECAY = 0.0001\nLR = 0.001\nMOMENTUM = 0\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, label, save=True):\n    save_dir = IMG_SAVE_DIR\n    save_name = f'{label}_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'Loss&IOU')\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        \ndef show_validation_score2(train_iou, save=True):\n    save_dir = IMG_SAVE_DIR\n    save_name = f'segmentation_validation_score.png'\n    \n    plt.plot(range(EPOCHS), train_iou, c='red', label='iou')\n    plt.xlabel('epoch')\n    plt.ylabel('iou')\n    plt.legend()\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    pixels = predicted_img.flatten()\n    # print(len(pixels))\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    # print(len(runs))\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\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，并将结果转换为标量值\n\n\ndef post_process(mask, min_size=10, shape=(520, 704,)):\n    num_component, component = cv2.connectedComponents(mask.astype(np.uint8))\n\n    predictions = []\n    for c in range(1, num_component):\n        p = (component == c)\n\n        if p.sum() > min_size:\n            a_prediction = np.zeros(shape, np.int64)\n            a_prediction[p] = 1\n            predictions.append(a_prediction)\n    contains_nonzero = np.any(predictions)\n    return predictions\n\n\ndef remove_isolated_points_from_rle(strin):\n    t2 = strin.split(\" \")\n    a = []\n    for i in range(0, len(t2), 2):\n        if t2[i + 1] != \"1\":\n            a.append(t2[i])\n            a.append(t2[i + 1])\n    return ' '.join(a)","metadata":{"execution":{"iopub.status.busy":"2023-12-27T11:49:36.531313Z","iopub.execute_input":"2023-12-27T11:49:36.531841Z","iopub.status.idle":"2023-12-27T11:49:36.591393Z","shell.execute_reply.started":"2023-12-27T11:49:36.531796Z","shell.execute_reply":"2023-12-27T11:49:36.590266Z"},"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        print(image_id)\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        print(f'shape: {np.shape(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-27T10:57:13.223016Z","iopub.execute_input":"2023-12-27T10:57:13.223835Z","iopub.status.idle":"2023-12-27T10:57:29.300268Z","shell.execute_reply.started":"2023-12-27T10:57:13.223795Z","shell.execute_reply":"2023-12-27T10:57:29.299234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# optimizers = ['SGD', 'Adam', 'AdamW']\n\ntrain_loss_list_total = []\nvalid_loss_list_total = []\nvalid_iou_list_total = []\n\nif RUN_TRAINING:\n    # 可视化用\n    train_loss_list = []\n    valid_loss_list = []\n    valid_iou_list = []\n\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        num_total = [p for p in range(len(grouped_df['id'].drop_duplicates()))]\n        np.random.seed(42)\n        np.random.shuffle(num_total)\n        split_point = int(0.8 * len(num_total))\n        trn_idx, val_idx = np.split(num_total, [split_point])\n        train_loader, valid_loader = create_dataloader(grouped_df, df, trn_idx, val_idx)\n        # 可视化用\n        valid_losses = []\n        valid_ious = []\n    train_ious = []\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_100).pth'))\n    criterion = nn.BCEWithLogitsLoss().to(device)\n\n    optimizer = optim.AdamW(model.parameters())\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        train_iou = 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            iou = IOU(output, masks)\n            loss = criterion(output, masks)\n            loss.backward()\n            optimizer.step()\n            # print(loss.item())\n            train_loss += loss.item()\n            train_iou += iou\n            \n        train_loss /= len(train_loader)\n        train_iou /= 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} - iou : {train_iou:.4f} / Exec time {exec_t} min\\n')\n            train_ious.append(train_iou)\n            train_losses.append(train_loss)\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 + 'pre_segmentation.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    if TRAIN_ALL == False:\n        show_validation_score(train_loss_list, valid_loss_list, valid_iou_list)\n    else:\n        show_validation_score2(train_ious)\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\nelse:\n    print('RUN_TRAINING is False')\n\nif 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(200).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                test_preds.append(predicted_mask)\n                predicted_mask = (predicted_mask > THRESHOLD).astype(np.uint8)\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\n    # print(submission_df)\n\n    # plt.imshow(predicted_mask > THRESHOLD)\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.title('Predict_TransUnet')\n    plt.show()\n    print('yes')\n\nelse:\n    print('RUN_INFERENCE is False')","metadata":{"execution":{"iopub.status.busy":"2023-12-27T11:49:40.286621Z","iopub.execute_input":"2023-12-27T11:49:40.286995Z","iopub.status.idle":"2023-12-27T11:50:14.056870Z","shell.execute_reply.started":"2023-12-27T11:49:40.286964Z","shell.execute_reply":"2023-12-27T11:50:14.055906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_preds = np.array(test_preds)\n# preds_test_t = (test_preds > THRESHOLD).astype(np.uint8)\n\n# predicted_nucleus = []\n# test_nucleus_image_id = []\n\n# for index, s in enumerate(preds_test_t):\n#     nucleus = post_process(cv2.resize(s, (704,520,), interpolation = cv2.INTER_LINEAR))\n#     for nucl in nucleus:\n#         predicted_nucleus.append(nucl)\n#         test_nucleus_image_id.append(ids[index])\n\n# predicted2 = [encode_rle(test_mask2) for test_mask2 in predicted_nucleus]\n# print(predicted2[0])\n# predicted_filt = [remove_isolated_points_from_rle(s) for s in predicted2]\n# print(predicted_filt[0])\n# submission_df = pd.DataFrame({'id':test_nucleus_image_id, 'predicted':predicted_filt})   \n# print(submission_df.head())","metadata":{"execution":{"iopub.status.busy":"2023-12-27T10:58:59.715220Z","iopub.execute_input":"2023-12-27T10:58:59.715533Z","iopub.status.idle":"2023-12-27T10:58:59.720129Z","shell.execute_reply.started":"2023-12-27T10:58:59.715504Z","shell.execute_reply":"2023-12-27T10:58:59.719161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# submission_df.to_csv('submission.csv', index=False)\n# target = submission_df.iloc[1]\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()","metadata":{"execution":{"iopub.status.busy":"2023-12-27T10:58:59.721545Z","iopub.execute_input":"2023-12-27T10:58:59.721844Z","iopub.status.idle":"2023-12-27T10:58:59.730087Z","shell.execute_reply.started":"2023-12-27T10:58:59.721809Z","shell.execute_reply":"2023-12-27T10:58:59.729307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# save_dir=IMG_SAVE_DIR\n# x = [p for p in range(1, EPOCHS + 1)]\n# for i in range(3):\n#     plt.plot(x, train_loss_list_total[i][0], label = f'Dropout = {DropOut[i]}')\n#     plt.title('Train_Loss')\n# plt.xlabel('Epoch')\n# plt.ylabel('Train_Loss')\n# plt.legend()\n# save_name = f'Train_Loss.png'\n# plt.savefig(save_dir + save_name)\n# plt.show()\n\n# for i in range(3):   \n#     plt.plot(x, valid_loss_list_total[i][0], label = f'Dropout = {DropOut[i]}')\n#     plt.title('Valid_Loss')\n# plt.xlabel('Epoch')\n# plt.ylabel('Valid_Loss')\n# plt.legend()\n# save_name = f'Valid_Loss.png'\n# plt.savefig(save_dir + save_name)\n# plt.show()\n\n# for i in range(3):\n#     plt.plot(x, valid_iou_list_total[i][0], label = f'Dropout = {DropOut[i]}')\n#     plt.title('Valid_IOU')\n# plt.xlabel('Epoch')\n# plt.ylabel('Valid_IOU')\n# plt.legend()\n# save_name = f'Valid_IOU.png'\n# plt.savefig(save_dir + save_name)\n# print('img save successfully!')\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-27T10:58:59.731308Z","iopub.execute_input":"2023-12-27T10:58:59.731636Z","iopub.status.idle":"2023-12-27T10:58:59.742100Z","shell.execute_reply.started":"2023-12-27T10:58:59.731603Z","shell.execute_reply":"2023-12-27T10:58:59.741328Z"},"trusted":true},"execution_count":null,"outputs":[]}]}