{"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":22990,"databundleVersionId":2048213,"sourceType":"competition"},{"sourceId":56486072,"sourceType":"kernelVersion"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"------------------------","metadata":{}},{"cell_type":"markdown","source":"I create train and inference notebook for Unet++ of [segmentation_models.pytorch](https://github.com/qubvel/segmentation_models.pytorch).\n\nI use only pytorch for framework.\n\nThe accuracy of the model is going to be tuned and improved in the future.\n\nI published this notebook for our reference as an example implementation using only pytorch.","metadata":{}},{"cell_type":"markdown","source":"I refered following two great notebooks for training and inference.\n\n- https://www.kaggle.com/iafoss/hubmap-pytorch-fast-ai-starter\n\n- https://www.kaggle.com/curiosity806/hubmap-use-catalyst-smp-and-albumentations\n\nAnd refered following great notebook for dice loss.\n\n- https://www.kaggle.com/bigironsphere/loss-function-library-keras-pytorch\n\nKindly upvote and appreciate the original work.","metadata":{}},{"cell_type":"markdown","source":"## Load libraries","metadata":{}},{"cell_type":"code","source":"!pip install git+https://github.com/qubvel/segmentation_models.pytorch","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:01.046732Z","iopub.execute_input":"2024-01-28T05:22:01.047668Z","iopub.status.idle":"2024-01-28T05:22:29.681155Z","shell.execute_reply.started":"2024-01-28T05:22:01.047631Z","shell.execute_reply":"2024-01-28T05:22:29.679835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nimport os\nimport random\nimport time\nimport warnings\nwarnings.simplefilter(\"ignore\")\n\n#import pdb\n#import zipfile\n#import pydicom\nfrom albumentations import *\nfrom albumentations.pytorch import ToTensorV2\nimport cv2\nimport albumentations as A\nfrom matplotlib import pyplot as plt\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image, ImageFilter\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import KFold\nimport tifffile as tiff\nimport torch,timm\nimport torch.backends.cudnn as cudnn\nimport torch.nn as nn\nfrom torch.nn import functional as F\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset, sampler\nfrom tqdm import tqdm_notebook as tqdm\n\n%matplotlib inline","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-01-28T05:22:29.683479Z","iopub.execute_input":"2024-01-28T05:22:29.683814Z","iopub.status.idle":"2024-01-28T05:22:35.291465Z","shell.execute_reply.started":"2024-01-28T05:22:29.683769Z","shell.execute_reply":"2024-01-28T05:22:35.290278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#assert 3==5 #学習するときにはここをコメントアウト","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:35.292980Z","iopub.execute_input":"2024-01-28T05:22:35.293892Z","iopub.status.idle":"2024-01-28T05:22:35.298532Z","shell.execute_reply.started":"2024-01-28T05:22:35.293849Z","shell.execute_reply":"2024-01-28T05:22:35.297446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"!mkdir ./masks\n!mkdir ./train\n\n!unzip ../input/256x256-images/masks.zip -d ./masks\n!unzip ../input/256x256-images/train.zip -d ./train","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-28T05:22:35.301566Z","iopub.execute_input":"2024-01-28T05:22:35.302354Z","iopub.status.idle":"2024-01-28T05:22:35.318339Z","shell.execute_reply.started":"2024-01-28T05:22:35.302212Z","shell.execute_reply":"2024-01-28T05:22:35.317209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"##  Set parameters","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=2**3):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(121)","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:35.319672Z","iopub.execute_input":"2024-01-28T05:22:35.320019Z","iopub.status.idle":"2024-01-28T05:22:35.336261Z","shell.execute_reply.started":"2024-01-28T05:22:35.319982Z","shell.execute_reply":"2024-01-28T05:22:35.335009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"To save time, the number of \"folds\" is smaller. \nThese days, depending on the model, I think it's more common to use 5 ~ 10.","metadata":{}},{"cell_type":"code","source":"fold = 0\nnfolds = 5\nreduce = 4\nsz = 256\n\nBATCH_SIZE = 2\nDEVICE = ('cuda' if torch.cuda.is_available() else 'cpu')\nEPOCHS = 30\nNUM_WORKERS = 4\nSEED = 2020\nTH = 0.39  #threshold for positive predictions\nmodel_name = \"Unet\"\nencoder_name = \"tu-\"+\"maxvit_rmlp_tiny_rw_256\"#\"maxvit_rmlp_nano_rw_256\" #https://smp.readthedocs.io/en/latest/encoders_timm.html \n#model_name=\"timmUnet\" #not work \n#encoder_name =  \"resnet34d\" \n#\"convnext_tiny_hnf.a2h_in1k\"\nencoder_lr = 1e-4\nif \"maxvit\" in encoder_name:\n    encoder_lr = (1e-4/8)*BATCH_SIZE #https://sorabatake.jp/33451/\ndecoder_lr = 1e-3\nmixup_P = 0.5 #mixupをする確率 #not work\ncutmix_P = 0\n\nif cutmix_P>0:\n    mixup_P = 0\n\nDATA = '../input/hubmap-kidney-segmentation/test/'\nLABELS = '../input/hubmap-kidney-segmentation/train.csv'\nMASKS = './masks/'\nTRAIN = './train/'\ndf_sample = pd.read_csv('../input/hubmap-kidney-segmentation/sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:35.337788Z","iopub.execute_input":"2024-01-28T05:22:35.338333Z","iopub.status.idle":"2024-01-28T05:22:35.384404Z","shell.execute_reply.started":"2024-01-28T05:22:35.338289Z","shell.execute_reply":"2024-01-28T05:22:35.383373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Util functions","metadata":{}},{"cell_type":"code","source":"#https://www.kaggle.com/bguberfain/memory-aware-rle-encoding\n#with bug fix\ndef rle_encode_less_memory(img):\n    #watch out for the bug\n    pixels = img.T.flatten()\n    \n    # This simplified method requires first and last pixel to be zero\n    pixels[0] = 0\n    pixels[-1] = 0\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    runs[1::2] -= runs[::2]\n    \n    return ' '.join(str(x) for x in runs)","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:35.386354Z","iopub.execute_input":"2024-01-28T05:22:35.386731Z","iopub.status.idle":"2024-01-28T05:22:35.393725Z","shell.execute_reply.started":"2024-01-28T05:22:35.386698Z","shell.execute_reply":"2024-01-28T05:22:35.392648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"# https://www.kaggle.com/iafoss/256x256-images\nmean = np.array([0.65527532, 0.49901106, 0.69247992] )\nstd = np.array([0.25565283, 0.31975344, 0.21533712])\n\ndef img2tensor(img,dtype:np.dtype=np.float32):\n    if img.ndim==2 : img = np.expand_dims(img,2)\n    img = np.transpose(img,(2,0,1))\n    return torch.from_numpy(img.astype(dtype, copy=False))\n\nclass HuBMAPDataset(Dataset):\n    def __init__(self, fold=fold, train=True, tfms=None):\n        ids = pd.read_csv(LABELS).id.values\n        kf = KFold(n_splits=nfolds,random_state=SEED,shuffle=True)\n        ids = set(ids[list(kf.split(ids))[fold][0 if train else 1]])\n        self.fnames = [fname for fname in os.listdir(TRAIN) if fname.split('_')[0] in ids]\n        self.train = train\n        self.tfms = tfms\n        \n    def __len__(self):\n        return len(self.fnames)\n    \n    def __getitem__(self, idx):\n        fname = self.fnames[idx]\n        img = cv2.cvtColor(cv2.imread(os.path.join(TRAIN,fname)), cv2.COLOR_BGR2RGB)\n        mask = cv2.imread(os.path.join(MASKS,fname),cv2.IMREAD_GRAYSCALE)\n        if self.tfms is not None:\n            augmented = self.tfms(image=img,mask=mask)\n            img,mask = augmented['image'],augmented['mask']\n        return img2tensor((img/255.0 - mean)/std),img2tensor(mask)","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:35.394971Z","iopub.execute_input":"2024-01-28T05:22:35.395299Z","iopub.status.idle":"2024-01-28T05:22:35.410398Z","shell.execute_reply.started":"2024-01-28T05:22:35.395247Z","shell.execute_reply":"2024-01-28T05:22:35.409320Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_aug(p=1.0):\n    return Compose([\n        HorizontalFlip(),\n        VerticalFlip(),\n        RandomRotate90(),\n        ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, \n                         border_mode=cv2.BORDER_REFLECT),\n        OneOf([\n            OpticalDistortion(p=0.3),\n            GridDistortion(p=.1),\n        ], p=0.3),\n        OneOf([\n            HueSaturationValue(10,15,10),\n            CLAHE(clip_limit=2),\n            RandomBrightnessContrast(),            \n        ], p=0.3),\n        #A.augmentations.geometric.resize.Resize(512,512)\n        #cutout,cutmix,mixup\n    ], p=p)","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:35.411745Z","iopub.execute_input":"2024-01-28T05:22:35.412111Z","iopub.status.idle":"2024-01-28T05:22:35.425444Z","shell.execute_reply.started":"2024-01-28T05:22:35.412075Z","shell.execute_reply":"2024-01-28T05:22:35.424460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#example of train images with masks\nds = HuBMAPDataset(tfms=get_aug())\ndl = DataLoader(ds,batch_size=64,shuffle=False,num_workers=NUM_WORKERS)\nimgs,masks = next(iter(dl))\n\nplt.figure(figsize=(16,16))\nfor i,(img,mask) in enumerate(zip(imgs,masks)):\n    img = ((img.permute(1,2,0)*std + mean)*255.0).numpy().astype(np.uint8)\n    plt.subplot(8,8,i+1)\n    plt.imshow(img,vmin=0,vmax=255)\n    plt.imshow(mask.squeeze().numpy(), alpha=0.2)\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\n    \ndel ds,dl,imgs,masks","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:35.429114Z","iopub.execute_input":"2024-01-28T05:22:35.429934Z","iopub.status.idle":"2024-01-28T05:22:43.838662Z","shell.execute_reply.started":"2024-01-28T05:22:35.429885Z","shell.execute_reply":"2024-01-28T05:22:43.836779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"from transformers import SegformerForSemanticSegmentation\n\n        \nclass Segformer(nn.Module):\n    def __init__(self, pretrained_model_name=\"nvidia/segformer-b3-finetuned-ade-512-512\"):#nvidia/segformer-b3-finetuned-ade-512-512 #nvidia/mit-b3\n        super(Segformer, self).__init__()\n        self.model = SegformerForSemanticSegmentation.from_pretrained(\n            pretrained_model_name,#かえてよい\n            id2label={0:\"glo\"},\n            label2id={\"glo\":0},ignore_mismatched_sizes=True\n        )\n        \n    def forward(self, inputs):\n        logits  = self.model(inputs).logits\n        #logits = nn.functional.interpolate(\n        #        logits, size=inputs.shape[-2:], mode=\"nearest\", align_corners=None\n        #    )\n        logits = nn.functional.interpolate(logits, size=inputs.shape[-2:], mode=\"bilinear\",align_corners=False).contiguous()\n        \n        \n        return logits\n    \n    \n#model = Segformer()\n#print(model(torch.randn((4,3,256,256))).shape)\n        \n        ","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:43.840706Z","iopub.execute_input":"2024-01-28T05:22:43.841230Z","iopub.status.idle":"2024-01-28T05:22:45.193085Z","shell.execute_reply.started":"2024-01-28T05:22:43.841172Z","shell.execute_reply":"2024-01-28T05:22:45.191643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ConvSilu(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size=3):\n        super(ConvSilu, self).__init__()\n        self.layer = nn.Sequential(\n            nn.Conv2d(in_channels=in_channels, out_channels=out_channels, kernel_size=kernel_size, padding=1),\n            nn.SiLU(inplace=True)\n        )\n    def forward(self, x):\n        return self.layer(x)\nclass Timm_Unet(nn.Module):\n    def __init__(self, name='resnet34', pretrained=True, inp_size=3, otp_size=1, decoder_filters=[32, 48, 64, 96, 128], **kwargs):\n        super(Timm_Unet, self).__init__()\n\n        if name.startswith('coat'):\n            print(\"Not yet\")\n            #encoder = coat_lite_medium()\n\n            #if pretrained:\n            #    checkpoint = './weights/coat_lite_medium_384x384_f9129688.pth'\n            #    checkpoint = torch.load(checkpoint, map_location=lambda storage, loc: storage)\n            #    state_dict = checkpoint['model']\n            #    encoder.load_state_dict(state_dict,strict=False)\n        \n            #encoder_filters = encoder.embed_dims\n        else:\n            encoder = timm.create_model(name, features_only=True, pretrained=pretrained, in_chans=inp_size)\n\n            encoder_filters = [f['num_chs'] for f in encoder.feature_info]\n\n        decoder_filters = decoder_filters\n\n        self.conv6 = ConvSilu(encoder_filters[-1], decoder_filters[-1])\n        self.conv6_2 = ConvSilu(decoder_filters[-1] + encoder_filters[-2], decoder_filters[-1])\n        self.conv7 = ConvSilu(decoder_filters[-1], decoder_filters[-2])\n        self.conv7_2 = ConvSilu(decoder_filters[-2] + encoder_filters[-3], decoder_filters[-2])\n        self.conv8 = ConvSilu(decoder_filters[-2], decoder_filters[-3])\n        self.conv8_2 = ConvSilu(decoder_filters[-3] + encoder_filters[-4], decoder_filters[-3])\n        self.conv9 = ConvSilu(decoder_filters[-3], decoder_filters[-4])\n\n        if len(encoder_filters) == 4:\n            self.conv9_2 = None\n        else:\n            self.conv9_2 = ConvSilu(decoder_filters[-4] + encoder_filters[-5], decoder_filters[-4])\n            \n        \n        self.conv10 = ConvSilu(decoder_filters[-4], decoder_filters[-5])\n        \n        self.res = nn.Conv2d(decoder_filters[-5], otp_size, 1, stride=1, padding=0)\n\n        self.cls =  nn.Linear(encoder_filters[-1] * 2, 5)\n        self.pix_sz =  nn.Linear(encoder_filters[-1] * 2, 1)\n\n        self._initialize_weights()\n\n        self.encoder = encoder\n\n\n    def forward(self, x):\n        batch_size, C, H, W = x.shape\n\n        if self.conv9_2 is None:\n            enc2, enc3, enc4, enc5 = self.encoder(x)\n        else:\n            enc1, enc2, enc3, enc4, enc5 = self.encoder(x)\n    \n    \n        dec6 = self.conv6(F.interpolate(enc5, scale_factor=2))\n        dec6 = self.conv6_2(torch.cat([dec6, enc4\n                ], 1))\n\n        dec7 = self.conv7(F.interpolate(dec6, scale_factor=2))\n        dec7 = self.conv7_2(torch.cat([dec7, enc3\n                ], 1))\n        \n        dec8 = self.conv8(F.interpolate(dec7, scale_factor=2))\n        dec8 = self.conv8_2(torch.cat([dec8, enc2\n                ], 1))\n\n        dec9 = self.conv9(F.interpolate(dec8, scale_factor=2))\n\n        if self.conv9_2 is not None:\n            dec9 = self.conv9_2(torch.cat([dec9, \n                    enc1\n                    ], 1))\n        \n        dec10 = self.conv10(dec9) # F.interpolate(dec9, scale_factor=2))\n\n        x1 = torch.cat([F.adaptive_avg_pool2d(enc5, output_size=1).view(batch_size, -1), \n                        F.adaptive_max_pool2d(enc5, output_size=1).view(batch_size, -1)], 1)\n\n        # x1 = F.dropout(x1, p=0.3, training=self.training)\n        #organ_cls = self.cls(x1)\n        #pixel_size = self.pix_sz(x1)\n        logits = self.res(dec10)\n        \n        logits = nn.functional.interpolate(logits, size=x.shape[-2:], mode=\"bilinear\",align_corners=False).contiguous()\n\n\n        return logits#, pixel_size\n\n\n    def _initialize_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d) or isinstance(m, nn.ConvTranspose2d) or isinstance(m, nn.Linear):\n                m.weight.data = nn.init.kaiming_normal_(m.weight.data)\n                if m.bias is not None:\n                    m.bias.data.zero_()\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:45.194668Z","iopub.execute_input":"2024-01-28T05:22:45.195259Z","iopub.status.idle":"2024-01-28T05:22:45.222443Z","shell.execute_reply.started":"2024-01-28T05:22:45.195226Z","shell.execute_reply":"2024-01-28T05:22:45.221343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder_name","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:45.224005Z","iopub.execute_input":"2024-01-28T05:22:45.224533Z","iopub.status.idle":"2024-01-28T05:22:45.244097Z","shell.execute_reply.started":"2024-01-28T05:22:45.224493Z","shell.execute_reply":"2024-01-28T05:22:45.243101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model():\n\n    #tf_efficientnet_b0,maxvit_nano_rw_256\n    #https://smp.readthedocs.io/en/latest/encoders_timm.htmlここからモデル選ぶ\n    #https://github.com/junkoda/kaggle_contrails_solution/tree/main/src/unet5 ぱらめーた参考に\n    \n    if model_name==\"Unet\":\n        model = smp.Unet(\n                     encoder_name=encoder_name,\n                     encoder_weights='imagenet',\n                     in_channels=3,\n                     classes=1)\n    elif model_name==\"UnetPlusPlus\":\n        model =  smp.UnetPlusPlus(\n                     encoder_name=encoder_name,\n                     encoder_weights='imagenet',\n                     in_channels=3,\n                     classes=1)\n    elif model_name==\"timmUnet\":#hubmap+HPA 2nd\n        model = Timm_Unet(name=encoder_name[3:])\n    elif model_name==\"segformer\":\n        model = Segformer()\n        \n    return model\n\n\nmodel = get_model()\n\n\nprint(model(torch.randn(4,3,256,256)).shape)\n","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:45.245499Z","iopub.execute_input":"2024-01-28T05:22:45.245916Z","iopub.status.idle":"2024-01-28T05:22:47.202576Z","shell.execute_reply.started":"2024-01-28T05:22:45.245878Z","shell.execute_reply":"2024-01-28T05:22:47.201539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## DiceLoss\n\nNote that this loss represents 1 - DiceLoss.","metadata":{}},{"cell_type":"code","source":"#https://www.kaggle.com/bigironsphere/loss-function-library-keras-pytorch\nclass DiceLoss(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(DiceLoss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        \n        #comment out if your model contains a sigmoid or equivalent activation layer\n        inputs = F.sigmoid(inputs)       \n        \n        #flatten label and prediction tensors\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()                            \n        dice = (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n        \n        return 1 - dice","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:47.204100Z","iopub.execute_input":"2024-01-28T05:22:47.204389Z","iopub.status.idle":"2024-01-28T05:22:47.212252Z","shell.execute_reply.started":"2024-01-28T05:22:47.204365Z","shell.execute_reply":"2024-01-28T05:22:47.211146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def mixup_data(x, y, alpha=1.0, use_cuda=True):\n    '''Returns mixed inputs, pairs of targets, and lambda'''\n    if alpha > 0:lam = np.random.beta(alpha, alpha)\n    else:lam = 1\n    batch_size = x.shape[0]#bs,seq_len,depth\n    if use_cuda:\n        index = torch.randperm(batch_size).cuda()\n    else:\n        index = torch.randperm(batch_size)\n    mixed_x = lam * x + (1 - lam) * x[index, :,:]\n    y_a, y_b = y, y[index]\n    return mixed_x, y_a, y_b, lam\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:47.213609Z","iopub.execute_input":"2024-01-28T05:22:47.213932Z","iopub.status.idle":"2024-01-28T05:22:47.224999Z","shell.execute_reply.started":"2024-01-28T05:22:47.213904Z","shell.execute_reply":"2024-01-28T05:22:47.224046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rand_bbox(size, lam):\n    W = size[0]\n    H = size[1]\n    cut_rat = np.sqrt(1. - lam)\n    cut_w = int(W * cut_rat)\n    cut_h = int(H * cut_rat)\n    cx = np.random.randint(W)\n    cy = np.random.randint(H)\n    bbx1 = np.clip(cx - cut_w // 2, 0, W)\n    bby1 = np.clip(cy - cut_h // 2, 0, H)\n    bbx2 = np.clip(cx + cut_w // 2, 0, W)\n    bby2 = np.clip(cy + cut_h // 2, 0, H)\n    return bbx1, bby1, bbx2, bby2\n\ndef cutmix_data(x, y, alpha=1.0, use_cuda=True):\n    '''Returns mixed inputs, pairs of targets, and lambda'''\n    if alpha > 0:lam = np.random.beta(alpha, alpha)\n    else:lam = 1\n    batch_size = x.shape[0]#bs,seq_len,depth\n    if use_cuda:\n        index = torch.randperm(batch_size).cuda()\n    else:\n        index = torch.randperm(batch_size)\n    bx1, by1, bx2, by2 = rand_bbox(x.shape[-2:], lam)\n    cutmix_img = x[index].clone()\n    cutmix_mask = y[index].clone()\n    cutmix_img = cutmix_img[:,:,bx1:bx2, by1:by2]\n    cutmix_mask = cutmix_mask[:,:,bx1:bx2, by1:by2]\n    x[:,:,bx1:bx2, by1:by2] = cutmix_img\n    y[:,:,bx1:bx2, by1:by2] = cutmix_mask\n    del cutmix_img,cutmix_mask\n    #y_a, y_b = y, y[index]\n    lam = 1 - ((bx2 - bx1) * (by2 - by1) / (x.shape[-1]*x.shape[-2]))\n    return x, y","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:47.226484Z","iopub.execute_input":"2024-01-28T05:22:47.226780Z","iopub.status.idle":"2024-01-28T05:22:47.242457Z","shell.execute_reply.started":"2024-01-28T05:22:47.226753Z","shell.execute_reply":"2024-01-28T05:22:47.241698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE=16","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:47.243512Z","iopub.execute_input":"2024-01-28T05:22:47.243774Z","iopub.status.idle":"2024-01-28T05:22:47.255421Z","shell.execute_reply.started":"2024-01-28T05:22:47.243751Z","shell.execute_reply":"2024-01-28T05:22:47.254400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder_lr*=4","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:47.256622Z","iopub.execute_input":"2024-01-28T05:22:47.256901Z","iopub.status.idle":"2024-01-28T05:22:47.267507Z","shell.execute_reply.started":"2024-01-28T05:22:47.256876Z","shell.execute_reply":"2024-01-28T05:22:47.266655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cv_score = 0\nfor fold in range(nfolds)[:1]:\n    ds_t = HuBMAPDataset(fold=fold, train=True, tfms=get_aug())\n    ds_v = HuBMAPDataset(fold=fold, train=False)\n    dataloader_t = torch.utils.data.DataLoader(ds_t,batch_size=BATCH_SIZE, shuffle=False,num_workers=NUM_WORKERS)\n    dataloader_v = torch.utils.data.DataLoader(ds_t,batch_size=BATCH_SIZE, shuffle=False,num_workers=NUM_WORKERS)\n    model = get_model().to(DEVICE)\n    \n    if model_name==\"segformer\":#時間かかる\n        encoder_lr = 2e-5*(BATCH_SIZE/8)#https://blog.roboflow.com/how-to-train-segformer-on-a-custom-dataset-with-pytorch-lightning/\n        optimizer = torch.optim.AdamW(model.parameters(),encoder_lr)#https://huggingface.co/blog/fine-tune-segformer \n    elif model_name==\"timmUnet\":\n        optimizer = torch.optim.AdamW(model.parameters(),encoder_lr)\n        \n    else:\n\n        optimizer = torch.optim.AdamW([\n            {'params': model.decoder.parameters(), 'lr': encoder_lr}, \n            {'params': model.encoder.parameters(), 'lr': decoder_lr},  \n        ])\n    scheduler = optim.lr_scheduler.OneCycleLR(optimizer=optimizer, pct_start=0.1, div_factor=1e3, \n                                              max_lr=encoder_lr*10, epochs=EPOCHS, steps_per_epoch=len(dataloader_t))\n    \n    diceloss = DiceLoss()\n    \n    print(f\"########FOLD: {fold}##############\")\n    \n    for epoch in range(EPOCHS):\n        ###Train\n        model.train()\n        train_loss = 0\n    \n        for data in tqdm(dataloader_t):\n            optimizer.zero_grad()\n            img, mask = data\n\n            img = img.to(DEVICE)\n            mask = mask.to(DEVICE)\n            rand = np.random.rand()\n            if mixup_P>rand:\n                img, y_a, y_b, lam = mixup_data(img, mask,alpha=0.5)\n            elif cutmix_P>rand:\n                img, mask = cutmix_data(img, mask,alpha=0.5)\n            outputs = model(img)\n    \n            #loss = diceloss(outputs, mask)\n            if mixup_P>rand:\n                loss = mixup_criterion(diceloss, outputs, y_a, y_b, lam)\n            elif cutmix_P>rand:\n                loss = diceloss(outputs, mask)\n            else:\n                loss = diceloss(outputs, mask)\n            loss.backward()\n            optimizer.step()\n            scheduler.step()\n            \n            train_loss += loss.item()\n        train_loss /= len(dataloader_t)\n        \n        print(f\"FOLD: {fold}, EPOCH: {epoch + 1}, train_loss: {train_loss}\")\n        \n        ###Validation\n        model.eval()\n        valid_loss = 0\n        \n        for data in dataloader_v:\n            img, mask = data\n            img = img.to(DEVICE)\n            mask = mask.to(DEVICE)\n        \n            outputs = model(img)\n    \n            loss = diceloss(outputs, mask)\n        \n            valid_loss += loss.item()\n        valid_loss /= len(dataloader_v)\n        \n        print(f\"FOLD: {fold}, EPOCH: {epoch + 1}, valid_loss: {valid_loss}\")\n        \n        \n    ###Save model\n    torch.save(model.state_dict(), f\"FOLD{fold}_.pth\")\n    \n    cv_score += valid_loss\n    \ncv_score = cv_score#/nfolds","metadata":{"execution":{"iopub.status.busy":"2024-01-28T05:22:47.268955Z","iopub.execute_input":"2024-01-28T05:22:47.269298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"CV score is: {1-cv_score}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -r ./masks\n!rm -r ./train","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}