{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\nbase_path = '/kaggle/input/hubmap-hacking-the-human-vasculature'\nannote_path = os.path.join(base_path,'polygons.jsonl')\nimgs_root_path = f'{base_path}/train'","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:42.063116Z","iopub.execute_input":"2023-08-08T06:45:42.063509Z","iopub.status.idle":"2023-08-08T06:45:42.068792Z","shell.execute_reply.started":"2023-08-08T06:45:42.063479Z","shell.execute_reply":"2023-08-08T06:45:42.067543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\n\ntiles_dicts = []\nwith open (annote_path,'r') as json_file:\n    #此时json_list为列表每一个索引存储了一个图片的信息\n     json_list = list(json_file)\n        \nfor json_str in json_list:\n    tiles_dicts.append(json.loads(json_str))","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:42.077328Z","iopub.execute_input":"2023-08-08T06:45:42.077603Z","iopub.status.idle":"2023-08-08T06:45:45.944913Z","shell.execute_reply.started":"2023-08-08T06:45:42.077579Z","shell.execute_reply":"2023-08-08T06:45:45.943823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\n# tiles_dicts[0]\nfor annot in tiles_dicts[0]['annotations']:\n    category = annot['type']\n    cords = annot['coordinates']\n    cords_array = np.array(cords,dtype = np.int32)\n    print(tiles_dicts[0]['id'])\n    break","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:45.947135Z","iopub.execute_input":"2023-08-08T06:45:45.947870Z","iopub.status.idle":"2023-08-08T06:45:45.959308Z","shell.execute_reply.started":"2023-08-08T06:45:45.947833Z","shell.execute_reply":"2023-08-08T06:45:45.958301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 制作分类标签","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport cv2\n#分类标签\ncategory_types = ['blood_vessel','glomerulus','unsure']\n\n'''\n    mask为np.ndarray\n    masks为list\n'''\n\nmasks = []\nids = []\nimgs_path = []\nfor idx in range(len(tiles_dicts)):\n    mask1 = np.zeros([512, 512, 1], np.uint8)\n    mask2 = np.zeros([512, 512, 1], np.uint8)\n    mask3 = np.zeros([512, 512, 1], np.uint8)\n    id = tiles_dicts[idx]['id']\n    ids.append(id)\n    imgs_path.append(f'{imgs_root_path}/{id}.tif')\n    for annot in tiles_dicts[idx]['annotations']:\n        category = annot['type']\n        cords = annot['coordinates']\n\n        cords_array = np.array(cords,dtype = np.int32)\n\n        if category == 'blood_vessel':\n                mask1 = cv2.fillPoly(mask1, [cords_array], 255)\n\n        elif category == 'glomerulus':\n                mask2 = cv2.fillPoly(mask2, [cords_array], 255)\n\n        elif category == 'unsure':\n                mask3 = cv2.fillPoly(mask3, [cords_array], 255)\n\n    mask = np.concatenate((mask1,mask2,mask3),axis = 2)\n    masks.append(mask)\n# plt.imshow(mask)\n# print(mask.shape)\n# print(len(masks))","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:45.960996Z","iopub.execute_input":"2023-08-08T06:45:45.961883Z","iopub.status.idle":"2023-08-08T06:45:49.413025Z","shell.execute_reply.started":"2023-08-08T06:45:45.961849Z","shell.execute_reply":"2023-08-08T06:45:49.412019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(masks))\n#存储了json中的所有id\nprint(len(ids)) \nprint(ids[0])\nprint(len(imgs_path))","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:49.415688Z","iopub.execute_input":"2023-08-08T06:45:49.416169Z","iopub.status.idle":"2023-08-08T06:45:49.421953Z","shell.execute_reply.started":"2023-08-08T06:45:49.416117Z","shell.execute_reply":"2023-08-08T06:45:49.420888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset与DataLoader","metadata":{}},{"cell_type":"code","source":"width = 512\nheight = 512","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:49.423507Z","iopub.execute_input":"2023-08-08T06:45:49.424183Z","iopub.status.idle":"2023-08-08T06:45:49.432348Z","shell.execute_reply.started":"2023-08-08T06:45:49.424131Z","shell.execute_reply":"2023-08-08T06:45:49.431291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset,DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\n\nclass MyDataset(Dataset):\n    \n    def __init__(self,masks,imgs_path,transforms = None):\n        \n        self.masks = masks\n        self.imgs_path = imgs_path\n        self.transforms = transforms\n            \n    def __getitem__(self,idx):\n        \n        img = Image.open(self.imgs_path[idx]).convert('RGB')\n        \n        mask = masks[idx]\n        \n        if self.transforms:\n            img  = self.transforms(img)\n            mask = self.transforms(mask)\n            \n#         return {'img':img,'mask':mask}\n            return img,mask\n        \n    def __len__(self):\n        return len(self.masks)\n    \ntransforms = transforms.Compose([\n#     transforms.Resize((width,height)),\n    transforms.ToTensor()\n])","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:49.433846Z","iopub.execute_input":"2023-08-08T06:45:49.434204Z","iopub.status.idle":"2023-08-08T06:45:49.443656Z","shell.execute_reply.started":"2023-08-08T06:45:49.434173Z","shell.execute_reply":"2023-08-08T06:45:49.442161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import random_split\ntotal_dataset = MyDataset(masks,imgs_path,transforms)\n\ntrain_size = int(0.7*len(total_dataset))\nval_size = len(total_dataset)-train_size\n\ntrain_dataset,val_dataset = random_split(total_dataset,[train_size,val_size])","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:49.445133Z","iopub.execute_input":"2023-08-08T06:45:49.445536Z","iopub.status.idle":"2023-08-08T06:45:49.457023Z","shell.execute_reply.started":"2023-08-08T06:45:49.445505Z","shell.execute_reply":"2023-08-08T06:45:49.455989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from torch.utils.data import random_split\n# total_dataset = MyDataset(masks,imgs_path,transforms)\n# # total_dataset[0]\n# print(total_dataset[0]['img'].shape)\n# print(total_dataset[0]['mask'].shape)\n\n\n# train_size = int(0.7*len(total_dataset))\n# val_size = len(total_dataset)-train_size\n\n# train_dataset,val_dataset = random_split(total_dataset,[train_size,val_size])\n# print(len(train_dataset))\n# print(len(val_dataset))\n# print(train_dataset[0]['img'].dtype)\n# print(train_dataset[0]['mask'].dtype)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:49.458544Z","iopub.execute_input":"2023-08-08T06:45:49.459006Z","iopub.status.idle":"2023-08-08T06:45:49.467460Z","shell.execute_reply.started":"2023-08-08T06:45:49.458975Z","shell.execute_reply":"2023-08-08T06:45:49.466463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\n#加载训练集与验证集\n#batch_size = 4，4个为一组\ntrain_dataloader = DataLoader(train_dataset,batch_size = 4,shuffle = True)\nval_dataloader = DataLoader(val_dataset,batch_size = 4,shuffle = True)\n\n# 获取一个 batch 的数据\ndata = next(iter(train_dataloader))\n\nprint(data[0].shape)\nprint(data[1].shape)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:49.468843Z","iopub.execute_input":"2023-08-08T06:45:49.469317Z","iopub.status.idle":"2023-08-08T06:45:49.541423Z","shell.execute_reply.started":"2023-08-08T06:45:49.469285Z","shell.execute_reply":"2023-08-08T06:45:49.540471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_mask(origin_tensor):\n    return torch.where(origin_tensor > 0.5, torch.tensor(1), torch.tensor(0))","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:49.544944Z","iopub.execute_input":"2023-08-08T06:45:49.545269Z","iopub.status.idle":"2023-08-08T06:45:49.551897Z","shell.execute_reply.started":"2023-08-08T06:45:49.545243Z","shell.execute_reply":"2023-08-08T06:45:49.549203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# UNet模型\n","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nimport torch\n\nclass DoubleConv(nn.Module):\n    def __init__(self,in_channels,out_channels):\n        super(DoubleConv,self).__init__()\n        self.conv = nn.Sequential(\n            #第一次卷积\n            nn.Conv2d(in_channels,out_channels,kernel_size = 3,padding = 1),\n            #一个批归一化层\n            nn.BatchNorm2d(out_channels),\n            #只取大于0的部分\n            nn.ReLU(inplace = True),\n            \n            nn.Conv2d(out_channels,out_channels,kernel_size = 3,padding = 1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace = True)\n        )\n        \n    def forward(self,x):\n        return self.conv(x)\n\nclass UNet(nn.Module):\n    def __init__(self,num_classes):\n        #初始化对象，方便forward中使用\n        super(UNet,self).__init__()\n        \n        #encode\n        self.encode1 = DoubleConv(3,64)\n        self.pooling1 = nn.MaxPool2d(kernel_size = 2,stride = 2)\n        \n        self.encode2 = DoubleConv(64,128)\n        self.pooling2 = nn.MaxPool2d(kernel_size = 2,stride = 2)\n        \n        self.encode3 = DoubleConv(128,256)\n        self.pooling3 = nn.MaxPool2d(kernel_size = 2,stride = 2)\n        \n        self.encode4 = DoubleConv(256,512)\n        self.pooling4 = nn.MaxPool2d(kernel_size = 2,stride = 2)\n        \n        self.mid = DoubleConv(512,1024)\n        \n        #decode\n        self.up1 = nn.ConvTranspose2d(1024,512,kernel_size = 2,stride = 2)\n        self.decode1 = DoubleConv(1024,512)#这里有copy and crop的部分\n        \n        \n        self.up2 = nn.ConvTranspose2d(512,256,kernel_size = 2,stride = 2)\n        self.decode2 = DoubleConv(512,256)\n        \n        self.up3 = nn.ConvTranspose2d(256,128,kernel_size = 2,stride = 2)\n        self.decode3 = DoubleConv(256,128)\n        \n        self.up4 = nn.ConvTranspose2d(128,64,kernel_size = 2,stride = 2)\n        self.decode4 = DoubleConv(128,64)\n        self.out = DoubleConv(64,num_classes)\n        \n        #使用softmax处理\n        self.softmax = torch.softmax\n        \n    def forward(self,x):\n        encode1 = self.encode1(x)\n        pooling1 = self.pooling1(encode1)\n        \n        encode2 = self.encode2(pooling1)\n        pooling2 = self.pooling2(encode2)\n        \n        encode3 = self.encode3(pooling2)\n        pooling3 = self.pooling3(encode3)\n        \n        encode4 = self.encode4(pooling3)\n        pooling4 = self.pooling4(encode4)\n        \n        mid = self.mid(pooling4)\n        #实现copy and crop\n        up1 = self.up1(mid)\n        up1 = torch.cat([up1,encode4],dim = 1)\n        decode1 = self.decode1(up1)\n        \n        up2 = self.up2(decode1)\n        up2 = torch.cat([up2,encode3],dim = 1)\n        decode2 = self.decode2(up2)\n        \n        up3 = self.up3(decode2)\n        up3 = torch.cat([up3,encode2],dim = 1)\n        decode3 = self.decode3(up3)\n        \n        up4 = self.up4(decode3)\n        up4 = torch.cat([up4,encode1],dim = 1)\n        decode4 = self.decode4(up4)\n        \n        out = self.out(decode4)\n        #使用softmax处理\n        out = self.softmax(out,dim = 1)\n        return out","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:49.553510Z","iopub.execute_input":"2023-08-08T06:45:49.553855Z","iopub.status.idle":"2023-08-08T06:45:49.575835Z","shell.execute_reply.started":"2023-08-08T06:45:49.553823Z","shell.execute_reply":"2023-08-08T06:45:49.574741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\ndef dice_coefficient(y_pred, y_true):\n    smooth = 1e-5  # 平滑因子，用于防止分母为零\n    intersection = torch.sum(y_pred * y_true)\n    union = torch.sum(y_pred) + torch.sum(y_true)\n    dice = (2.0 * intersection + smooth) / (union + smooth)\n    return dice\n\nclass DiceLoss(nn.Module):\n    def __init__(self):\n        super(DiceLoss, self).__init__()\n\n    def forward(self, y_pred, y_true):\n        dice = dice_coefficient(y_pred, y_true)\n        dice_loss = 1.0 - dice  # 将Dice系数转换为Dice loss\n        return dice_loss\n","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:49.577957Z","iopub.execute_input":"2023-08-08T06:45:49.578994Z","iopub.status.idle":"2023-08-08T06:45:49.589116Z","shell.execute_reply.started":"2023-08-08T06:45:49.578960Z","shell.execute_reply":"2023-08-08T06:45:49.588176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.optim.lr_scheduler as lr_scheduler\nimport torch.optim as optim\n\n# 创建UNet模型实例\nmodel = UNet(3)\n\n# 定义损失函数和优化器\ncriterion = DiceLoss()\noptimizer = optim.Adam(model.parameters())\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nmodel = model.float()\nmodel.to(device)\n\n# 开始训练\nnum_epochs = 60\ntrain_model_from_begining = True\n\n# 定义学习率调整器\nstep_size = 10  # 每隔5个epoch降低一次学习率\ngamma = 0.5  # 学习率降低的倍数\nscheduler = lr_scheduler.StepLR(optimizer, step_size=step_size, gamma=gamma)\n\nif train_model_from_begining:\n\n    for epoch in range(num_epochs):\n        running_loss = 0.0\n\n        for i, (inputs, labels) in enumerate(train_dataloader):\n    #         print(np.shape(inputs))\n\n            inputs = inputs.to(torch.float32)\n            labels = labels.to(torch.float32)\n\n            inputs = inputs.to(device)\n            labels = labels.to(device)\n\n            # 清除梯度\n            optimizer.zero_grad()\n\n            # 前向传播\n            outputs = model(inputs)\n    #         outputs = get_mask(outputs)\n\n            # 计算损失\n            loss = criterion(outputs, labels)\n\n            # 反向传播和优化\n            loss.requires_grad_(True)\n            loss.backward()\n            optimizer.step()\n\n            # 累积损失\n            running_loss += loss.item()\n\n        # 输出每个epoch的平均损失\n        print(f\"Epoch {epoch+1} - Loss: {running_loss/len(train_dataloader)}\")\n        \n        #更新学习率\n        scheduler.step()","metadata":{"execution":{"iopub.status.busy":"2023-08-08T06:45:49.590646Z","iopub.execute_input":"2023-08-08T06:45:49.591039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train_model_from_begining:\n    torch.save(model,\"model.pth\")\nelse:\n    model = torch.load(\"/kaggle/input/unetmodel/model_10epoch.pth\") ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}