{"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":"markdown","source":"![DIS](https://github.com/xuebinqin/DIS/raw/c5da300d44bb022c11a3752dabf0d88fa2b00e5d/figures/dis-logo-official.png)\n\n## **DIS Description**\nCurrently, existing image segmentation tasks mainly focus on segmenting objects with specific characteristics, e.g., salient, camouflaged, meticulous, or specific categories. Most of them have the same input/output formats, and barely use exclusive mechanisms designed for segmenting targets in their models, which means almost all tasks are dataset-dependent. Thus, it is very promising to formulate a category-agnostic DIS task for accurately segmenting objects with different structure complexities, regardless of their characteristics. Compared with semantic segmentation, the proposed DIS task usually focuses on images with single or a few targets, from which getting richer accurate details of each target is more feasible.\n\n**[GitHub Repo](https://github.com/xuebinqin/DIS)**\n\n**[Research Paper](https://arxiv.org/pdf/2203.03041.pdf)**","metadata":{}},{"cell_type":"code","source":"!ls ../input/hubmap-hpa-dichotomous-image-segmentation/DIS/IS-Net/saved_models/","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:36:44.706516Z","iopub.execute_input":"2022-09-16T03:36:44.707021Z","iopub.status.idle":"2022-09-16T03:36:45.967626Z","shell.execute_reply.started":"2022-09-16T03:36:44.706913Z","shell.execute_reply":"2022-09-16T03:36:45.966460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q segmentation_models_pytorch","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:36:45.970088Z","iopub.execute_input":"2022-09-16T03:36:45.970755Z","iopub.status.idle":"2022-09-16T03:37:01.073409Z","shell.execute_reply.started":"2022-09-16T03:36:45.970716Z","shell.execute_reply":"2022-09-16T03:37:01.072160Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!git clone https://github.com/xuebinqin/DIS.git","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:01.076534Z","iopub.execute_input":"2022-09-16T03:37:01.081125Z","iopub.status.idle":"2022-09-16T03:37:05.917172Z","shell.execute_reply.started":"2022-09-16T03:37:01.081083Z","shell.execute_reply":"2022-09-16T03:37:05.916008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%cd ./DIS/IS-Net\n\n!pip install gdown","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:05.920272Z","iopub.execute_input":"2022-09-16T03:37:05.920893Z","iopub.status.idle":"2022-09-16T03:37:29.045775Z","shell.execute_reply.started":"2022-09-16T03:37:05.920836Z","shell.execute_reply":"2022-09-16T03:37:29.044614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:29.047647Z","iopub.execute_input":"2022-09-16T03:37:29.048350Z","iopub.status.idle":"2022-09-16T03:37:29.053600Z","shell.execute_reply.started":"2022-09-16T03:37:29.048307Z","shell.execute_reply":"2022-09-16T03:37:29.052622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Imports**","metadata":{}},{"cell_type":"code","source":"\nimport os\nimport glob\n\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\nfrom albumentations import *\nfrom sklearn.model_selection import KFold\nfrom PIL import Image\nimport tifffile as tiff \nimport random\nfrom glob import glob\nimport os, shutil\nfrom tqdm import tqdm\ntqdm.pandas()\nimport time\nimport copy\nimport joblib\nfrom collections import defaultdict\nimport gc\nimport gdown\nimport requests\nfrom io import BytesIO\nfrom IPython import display as ipd\n\nimport torch \nimport torch.nn as nn\nfrom torch.autograd import Variable\nfrom torchvision import transforms\nimport torchvision.transforms.functional as TF\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda import amp\nimport torch.nn.functional as F\n\nimport segmentation_models_pytorch as smp\n\nfrom models import *            # { loading ISNet segementation model   }\nfrom data_loader_cache import * # | data loaders utilities              |\nfrom basics import  *           # { & evaluation functions from DIS repo}    \n\nfrom colorama import Fore, Back, Style\nc_  = Fore.GREEN\nsr_ = Style.RESET_ALL","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:29.054961Z","iopub.execute_input":"2022-09-16T03:37:29.058547Z","iopub.status.idle":"2022-09-16T03:37:39.979071Z","shell.execute_reply.started":"2022-09-16T03:37:29.058451Z","shell.execute_reply":"2022-09-16T03:37:39.978040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**folder paths**","metadata":{}},{"cell_type":"code","source":"TRAIN = '../../../input/hubmap-2022-512x512/train/'\nMASKS = '../../../input/hubmap-2022-512x512/masks'\ntrain_csv = '../../../input/hubmap-organ-segmentation/train.csv'","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:39.980931Z","iopub.execute_input":"2022-09-16T03:37:39.981651Z","iopub.status.idle":"2022-09-16T03:37:39.988434Z","shell.execute_reply.started":"2022-09-16T03:37:39.981609Z","shell.execute_reply":"2022-09-16T03:37:39.987508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pwd","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:39.990153Z","iopub.execute_input":"2022-09-16T03:37:39.990944Z","iopub.status.idle":"2022-09-16T03:37:40.997535Z","shell.execute_reply.started":"2022-09-16T03:37:39.990908Z","shell.execute_reply":"2022-09-16T03:37:40.996403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Image visualization**","metadata":{}},{"cell_type":"code","source":"bs = 64\nnfolds = 4\nfold = 3\nSEED = 2022\nNUM_WORKERS = 2\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:40.999663Z","iopub.execute_input":"2022-09-16T03:37:41.000055Z","iopub.status.idle":"2022-09-16T03:37:41.071337Z","shell.execute_reply.started":"2022-09-16T03:37:41.000015Z","shell.execute_reply":"2022-09-16T03:37:41.070232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndf = pd.read_csv(train_csv)\n\ndf.head(10)","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:41.076410Z","iopub.execute_input":"2022-09-16T03:37:41.077356Z","iopub.status.idle":"2022-09-16T03:37:41.374165Z","shell.execute_reply.started":"2022-09-16T03:37:41.077317Z","shell.execute_reply":"2022-09-16T03:37:41.373253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# img_1 = Image.open(\"../input/hubmap-organ-segmentation/train_images/14756.tiff\")\n# img_1.show()\nimg_id_1 = 10274\nimg_1 = tiff.imread(f\"../../../input/hubmap-organ-segmentation/train_images/{img_id_1}.tiff\")\nprint(img_1.shape)\nprint(type(img_1))","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:41.375690Z","iopub.execute_input":"2022-09-16T03:37:41.376560Z","iopub.status.idle":"2022-09-16T03:37:41.749931Z","shell.execute_reply.started":"2022-09-16T03:37:41.376521Z","shell.execute_reply":"2022-09-16T03:37:41.748907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15,15))\nplt.imshow(img_1)\nplt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:41.752176Z","iopub.execute_input":"2022-09-16T03:37:41.753280Z","iopub.status.idle":"2022-09-16T03:37:43.329580Z","shell.execute_reply.started":"2022-09-16T03:37:41.753236Z","shell.execute_reply":"2022-09-16T03:37:43.328289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Loading Masks**","metadata":{}},{"cell_type":"code","source":"\n# https://www.kaggle.com/paulorzp/rle-functions-run-length-encode-decode\ndef mask2rle(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels= img.T.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n \ndef rle2mask(mask_rle, shape=(1600,256)):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (width,height) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n\n    '''\n    s = mask_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(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape).T\n\n","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:43.330991Z","iopub.execute_input":"2022-09-16T03:37:43.331405Z","iopub.status.idle":"2022-09-16T03:37:43.344644Z","shell.execute_reply.started":"2022-09-16T03:37:43.331370Z","shell.execute_reply":"2022-09-16T03:37:43.341572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask_1 = rle2mask(df[df[\"id\"]==img_id_1][\"rle\"].iloc[-1], (img_1.shape[1], img_1.shape[0]))\nmask_1.shape","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:43.346062Z","iopub.execute_input":"2022-09-16T03:37:43.346628Z","iopub.status.idle":"2022-09-16T03:37:43.389331Z","shell.execute_reply.started":"2022-09-16T03:37:43.346594Z","shell.execute_reply":"2022-09-16T03:37:43.388339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15,15))\nplt.imshow(mask_1, cmap='coolwarm', alpha=0.5)\nplt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:43.390863Z","iopub.execute_input":"2022-09-16T03:37:43.391537Z","iopub.status.idle":"2022-09-16T03:37:44.492174Z","shell.execute_reply.started":"2022-09-16T03:37:43.391498Z","shell.execute_reply":"2022-09-16T03:37:44.491258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Combining Masks with Images**","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(15,15))\nplt.imshow(img_1)\nplt.imshow(mask_1, cmap='coolwarm', alpha=0.5)\nplt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:44.493786Z","iopub.execute_input":"2022-09-16T03:37:44.494171Z","iopub.status.idle":"2022-09-16T03:37:47.104066Z","shell.execute_reply.started":"2022-09-16T03:37:44.494133Z","shell.execute_reply":"2022-09-16T03:37:47.103219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Data**","metadata":{}},{"cell_type":"markdown","source":"### Pre-computed Dataset: \n[Dataset (512 x 512)](https://www.kaggle.com/datasets/thedevastator/hubmap-2022-512x512/)\n\n[Dataset (256 x 256)](https://www.kaggle.com/datasets/thedevastator/hubmap-2022-256x256/)\n\nby [@thedevastator](https://www.kaggle.com/thedevastator)","metadata":{}},{"cell_type":"code","source":"#dataset info\ndataset_HuBMAP_HPA =   {\"name\": \"HuBMAP_HPA_512x512\",\n                         \"im_dir\": \"../../../input/hubmap-2022-512x512/train\",\n                         \"gt_dir\": \"../../../input/hubmap-2022-512x512/masks\",\n                         \"im_ext\": \".png\",\n                         \"gt_ext\": \".png\",\n                         \"cache_dir\":\"cache_dir\"}","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:47.105606Z","iopub.execute_input":"2022-09-16T03:37:47.106201Z","iopub.status.idle":"2022-09-16T03:37:47.111745Z","shell.execute_reply.started":"2022-09-16T03:37:47.106161Z","shell.execute_reply":"2022-09-16T03:37:47.110927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Get image paths of dataset splits in a dict \ndef get_im_gt_name_dict(dataset, train = True):\n    print(\"------------------------------\", \"Train\" if train else \"Val\", \"--------------------------------\")\n    im_gt_list = []\n    \n    for i in range(len(dataset)):\n        \n        ids = pd.read_csv(train_csv).id.astype(str).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        \n        tmp_im_list = [os.path.join(dataset[i][\"im_dir\"],fname) for fname in os.listdir(dataset[i][\"im_dir\"]) if fname.split('_')[0] in ids]\n        tmp_gt_list = [os.path.join(dataset[i][\"gt_dir\"],fname) for fname in os.listdir(dataset[i][\"gt_dir\"]) if fname.split('_')[0] in ids]\n        \n        print('-im-',dataset[i][\"name\"],dataset[i][\"im_dir\"], ': ',len(tmp_im_list))\n        print('-gt-', dataset[i][\"name\"],dataset[i][\"gt_dir\"], ': ',len(tmp_gt_list))\n        \n        if train:\n            cache_dir = \"cache_dir\"\n        else:\n            cache_dir = \"cache_dir_val\"\n        \n        im_gt_list.append({\"dataset_name\":dataset[i][\"name\"],\n                                        \"im_path\":tmp_im_list,\n                                        \"gt_path\":tmp_gt_list,\n                                        \"im_ext\":dataset[i][\"im_ext\"],\n                                        \"gt_ext\":dataset[i][\"gt_ext\"],\n                                        \"cache_dir\":cache_dir})\n        \n        return im_gt_list","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:47.113314Z","iopub.execute_input":"2022-09-16T03:37:47.114067Z","iopub.status.idle":"2022-09-16T03:37:47.128995Z","shell.execute_reply.started":"2022-09-16T03:37:47.114029Z","shell.execute_reply":"2022-09-16T03:37:47.128115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = [dataset_HuBMAP_HPA]\n\ntrain_nm_im_gt_list = get_im_gt_name_dict(dataset, train=True)\nvalid_nm_im_gt_list = get_im_gt_name_dict(dataset, train=False)","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:47.130489Z","iopub.execute_input":"2022-09-16T03:37:47.131073Z","iopub.status.idle":"2022-09-16T03:37:47.862054Z","shell.execute_reply.started":"2022-09-16T03:37:47.131040Z","shell.execute_reply":"2022-09-16T03:37:47.860956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_nm_im_gt_list[0].keys())\nprint(train_nm_im_gt_list[0][\"im_path\"][1001])\nprint(train_nm_im_gt_list[0][\"gt_path\"][1001])\nprint(train_nm_im_gt_list[0][\"cache_dir\"])\n\n\nprint(valid_nm_im_gt_list[0][\"im_path\"][100])\nprint(valid_nm_im_gt_list[0][\"gt_path\"][100])\nprint(valid_nm_im_gt_list[0][\"cache_dir\"])","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:47.863714Z","iopub.execute_input":"2022-09-16T03:37:47.864068Z","iopub.status.idle":"2022-09-16T03:37:47.873236Z","shell.execute_reply.started":"2022-09-16T03:37:47.864031Z","shell.execute_reply":"2022-09-16T03:37:47.871137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Image Tiles Visualization**","metadata":{}},{"cell_type":"code","source":"mean = np.array([0.7720342, 0.74582646, 0.76392896])\nstd = np.array([0.24745085, 0.26182273, 0.25782376])\n\nmean_n = np.array([[[0.7720342]], [[0.74582646]], [[0.76392896]]])\nstd_n = np.array([[[0.24745085]], [[0.26182273]], [[0.25782376]]])\n\nmean_m = np.array([[[0.7720342], [0.74582646], [0.76392896]]])\nstd_m  = np.array([[[0.24745085], [0.26182273], [0.25782376]]])","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:47.875200Z","iopub.execute_input":"2022-09-16T03:37:47.876476Z","iopub.status.idle":"2022-09-16T03:37:47.884235Z","shell.execute_reply.started":"2022-09-16T03:37:47.876441Z","shell.execute_reply":"2022-09-16T03:37:47.883240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(mean_n.shape)\nprint(std_n.shape)","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:47.885743Z","iopub.execute_input":"2022-09-16T03:37:47.886259Z","iopub.status.idle":"2022-09-16T03:37:47.899108Z","shell.execute_reply.started":"2022-09-16T03:37:47.886221Z","shell.execute_reply":"2022-09-16T03:37:47.898052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\ndef img2tensor(img,dtype:np.dtype=np.float32):\n    if img.ndim==2 : img = np.expand_dims(img,2)\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(train_csv).id.astype(str).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#         img = np.moveaxis(img,-1,0).astype('float64')\n#         img = img/255.0\n        \n        mask = cv2.imread(os.path.join(MASKS,fname),cv2.IMREAD_GRAYSCALE)\n#         mask = np.moveaxis(mask,-1,0).astype('float64')\n        \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),img2tensor(mask)\n#         return img2tensor((img - mean)/std),img2tensor(mask)\n    \n    \ndef 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            ElasticTransform(p=.3),\n            GaussianBlur(p=.3),\n            GaussNoise(p=.3),\n            OpticalDistortion(p=0.3),\n            GridDistortion(p=.1),\n            IAAPiecewiseAffine(p=0.3),\n        ], p=0.3),\n        RandomBrightnessContrast(brightness_limit=[-0.3,-0.1], contrast_limit=[0.1,0.3], brightness_by_max=False,p=0.6),\n        OneOf([\n            HueSaturationValue(15,25,0),\n            CLAHE(clip_limit=2),            \n        ], p=0.6),\n        Normalize (mean=mean, std=std, max_pixel_value=255.0, always_apply=True, p=1.0)\n    ], p=p)\n","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:47.901459Z","iopub.execute_input":"2022-09-16T03:37:47.902040Z","iopub.status.idle":"2022-09-16T03:37:47.917061Z","shell.execute_reply.started":"2022-09-16T03:37:47.902012Z","shell.execute_reply":"2022-09-16T03:37:47.916095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n#example of train images with masks\nds = HuBMAPDataset(tfms=get_aug())\ndl = DataLoader(ds,batch_size=32,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*std + mean)*255.0).numpy().astype(np.uint8)\n#     img = np.moveaxis(img,0,-1 )\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":"2022-09-16T03:37:47.918479Z","iopub.execute_input":"2022-09-16T03:37:47.919615Z","iopub.status.idle":"2022-09-16T03:37:57.575451Z","shell.execute_reply.started":"2022-09-16T03:37:47.919578Z","shell.execute_reply":"2022-09-16T03:37:57.574547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Set DIS Parameters** ","metadata":{}},{"cell_type":"code","source":"hypar = {} # paramters for inferencing\n\n\nhypar[\"model_path\"] =\"./saved_models\" ## load trained weights from this path\nhypar[\"restore_model\"] = \"../../../input/hubmap-hpa-dichotomous-image-segmentation/DIS/IS-Net/saved_models/ISNet_traLoss_2.1562_valLoss_0.6767_maxF1_0.4898_mae_0.0645.pth\" ## name of the to-be-loaded weights\"\nhypar[\"gt_encoder_model\"] = \"../../../input/hubmap-hpa-dichotomous-image-segmentation/DIS/IS-Net/saved_models/GTENCODER13700_traLoss_0.076_valLoss_0.0585_maxF1_0.1677_mae_0.1212.pth\"\n\nhypar[\"interm_sup\"] = True ## indicate if activate intermediate feature supervision\nhypar[\"valid_out_dir\"] = \"\"\n\n##  choose floating point accuracy --\nhypar[\"model_digit\"] = \"full\" ## indicates \"half\" or \"full\" accuracy of float number\nhypar[\"seed\"] = 0\nhypar[\"mode\"] =\"train\"\n\nhypar[\"cache_size\"] = [512, 512] ## cached input spatial resolution, can be configured into different size\nhypar[\"cache_boost_train\"] = False\nhypar[\"cache_boost_valid\"] = False\n## data augmentation parameters ---\nhypar[\"input_size\"] = [512, 512] ## mdoel input spatial size, usually use the same value hypar[\"cache_size\"], which means we don't further resize the images\nhypar[\"crop_size\"] = [512, 512] ## random crop size from the input, it is usually set as smaller than hypar[\"cache_size\"], e.g., [920,920] for data augmentation\n\n# print(\"building model...\")\nhypar[\"model\"] = ISNetDIS() #U2NETFASTFEATURESUP()\nhypar[\"early_stop\"] = 20 ## stop the training when no improvement in the past 20 validation periods, smaller numbers can be used here e.g., 5 or 10.\nhypar[\"model_save_fre\"] = 2000 ## valid and save model weights every 4000 iterations\nhypar[\"start_ite\"] = 0\n\nhypar[\"batch_size_train\"] = 8 ## batch size for training\nhypar[\"batch_size_valid\"] = 1 ## batch size for validation and inferencing\n# print(\"batch size: \", hypar[\"batch_size_train\"])\n\nhypar[\"max_ite\"] = 1000000## if early stop couldn't stop the training process, stop it by the max_ite_num\nhypar[\"max_epoch_num_gte\"] = 50 #number of epochs to train gt encoder\nhypar[\"max_epoch_num\"] = 150","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:45:08.978636Z","iopub.execute_input":"2022-09-16T03:45:08.979349Z","iopub.status.idle":"2022-09-16T03:45:09.357321Z","shell.execute_reply.started":"2022-09-16T03:45:08.979312Z","shell.execute_reply":"2022-09-16T03:45:09.356164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Dataloaders**","metadata":{}},{"cell_type":"code","source":"class GOSDatasetCache(GOSDatasetCache):\n    \n    def __getitem__(self, idx):\n        \n        im = None\n        gt = None\n        if(self.cache_boost and self.ims_pt is not None):\n\n            # start = time.time()\n            im = self.ims_pt[idx]#.type(torch.float32)\n            gt = self.gts_pt[idx]#.type(torch.float32)\n            # print(idx, 'time for pt loading: ', time.time()-start)\n\n        else:\n            # import time\n            # start = time.time()\n            # print(\"tensor***\")\n            im_pt_path = os.path.join(self.cache_path,os.sep.join(self.dataset[\"im_path\"][idx].split(os.sep)[-2:]))\n            im = torch.load(im_pt_path)#(self.dataset[\"im_path\"][idx])\n            im = torch.moveaxis(im,0,-1)\n            gt_pt_path = os.path.join(self.cache_path,os.sep.join(self.dataset[\"gt_path\"][idx].split(os.sep)[-2:]))\n            gt = torch.load(gt_pt_path)#(self.dataset[\"gt_path\"][idx])\n            gt = torch.moveaxis(gt,0,-1)\n            # print(idx,'time for tensor loading: ', time.time()-start)\n\n\n        im_shp =  im.shape #self.dataset[\"im_shp\"][idx]\n        # print(\"time for loading im and gt: \", time.time()-start)\n\n        # start_time = time.time()\n#         im = torch.divide(im,255.0)\n#         gt = torch.divide(gt,255.0)\n        # print(idx, 'time for normalize torch divide: ', time.time()-start_time)\n\n        sample = {\n        \"imidx\": torch.from_numpy(np.array(idx)),\n        \"image\": im.numpy(),\n        \"label\": gt.numpy(),\n        \"shape\": torch.from_numpy(np.array(im_shp)),\n        }\n\n        if self.transform:\n#             sample = self.transform(sample)\n            augmented = self.transform(image=sample[\"image\"],mask=sample[\"label\"])\n            sample[\"image\"] = torch.moveaxis(torch.from_numpy(augmented['image']),-1,0)\n            sample[\"label\"] = torch.moveaxis(torch.from_numpy(augmented['mask']),-1, 0)\n            \n        return sample\n\n    \ndef create_dataloaders(name_im_gt_list, cache_size=[], cache_boost=True, my_transforms=[], batch_size=1, shuffle=False):\n    ## model=\"train\": return one dataloader for training\n    ## model=\"valid\": return a list of dataloaders for validation or testing\n\n    gos_dataloaders = []\n    gos_datasets = []\n\n    if(len(name_im_gt_list)==0):\n        return gos_dataloaders, gos_datasets\n\n    num_workers_ = 1\n    if(batch_size>1):\n        num_workers_ = 2\n    if(batch_size>4):\n        num_workers_ = 4\n    if(batch_size>8):\n        num_workers_ = 8\n\n    for i in range(0,len(name_im_gt_list)):\n        gos_dataset = GOSDatasetCache([name_im_gt_list[i]],\n                                      cache_size = cache_size,\n                                      cache_path = name_im_gt_list[i][\"cache_dir\"],\n                                      cache_boost = cache_boost,\n                                      transform = my_transforms)\n        gos_dataloaders.append(DataLoader(gos_dataset, batch_size=batch_size, shuffle=shuffle, num_workers=num_workers_))\n        gos_datasets.append(gos_dataset)\n\n    return gos_dataloaders, gos_datasets","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:37:58.032712Z","iopub.execute_input":"2022-09-16T03:37:58.033089Z","iopub.status.idle":"2022-09-16T03:37:58.049672Z","shell.execute_reply.started":"2022-09-16T03:37:58.033050Z","shell.execute_reply":"2022-09-16T03:37:58.048546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloaders, train_datasets = create_dataloaders(train_nm_im_gt_list,\n                                                             cache_size = hypar[\"cache_size\"],\n                                                             cache_boost = hypar[\"cache_boost_train\"],\n                                                             my_transforms = get_aug(p=1.0),\n                                                             batch_size = hypar[\"batch_size_train\"],\n                                                             shuffle = True)\ntrain_dataloaders_val, train_datasets_val = create_dataloaders(train_nm_im_gt_list,\n                                                     cache_size = hypar[\"cache_size\"],\n                                                     cache_boost = hypar[\"cache_boost_train\"],\n                                                     my_transforms = get_aug(p=1.0),\n                                                     batch_size = hypar[\"batch_size_valid\"],\n                                                     shuffle = False)\n","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"--- create valid dataloader ---\")\n## build dataloader for validation or testing\n# valid_nm_im_gt_list = get_im_gt_name_dict(valid_datasets, flag=\"valid\")\n## build dataloader for training datasets\nvalid_dataloaders, valid_datasets = create_dataloaders(valid_nm_im_gt_list,\n                                                      cache_size = hypar[\"cache_size\"],\n                                                      cache_boost = hypar[\"cache_boost_valid\"],\n                                                      my_transforms = Compose([Normalize (mean=mean, std=std, max_pixel_value=255.0, always_apply=True, p=1.0)],p=1.0),\n                                                      batch_size=hypar[\"batch_size_valid\"],\n                                                      shuffle=False)\n","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(train_dataloaders[0]), \" train dataloaders created\")\nprint(\"length of the train dataloaders :\", len(train_dataloaders[0]))\n\nprint(len(valid_dataloaders), \" valid dataloaders created\")\nprint(\"length of the valid dataloaders :\", len(valid_dataloaders[0]))","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:40:02.522888Z","iopub.execute_input":"2022-09-16T03:40:02.523292Z","iopub.status.idle":"2022-09-16T03:40:02.529771Z","shell.execute_reply.started":"2022-09-16T03:40:02.523252Z","shell.execute_reply":"2022-09-16T03:40:02.528804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(train_datasets), \" train dataset created\")\nprint(len(valid_datasets), \" valid dataset created\")","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:40:02.531233Z","iopub.execute_input":"2022-09-16T03:40:02.532191Z","iopub.status.idle":"2022-09-16T03:40:02.544675Z","shell.execute_reply.started":"2022-09-16T03:40:02.532154Z","shell.execute_reply":"2022-09-16T03:40:02.543609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **IS-Net Model Architecture**\n**IS-Net** consists of a ***ground truth (GT) encoder***, a ***image segmentation component***, and a newly proposed ***intermediate supervision strategy***. The GT encoder (27.7 MB) is designed to encode the GT masks into high-dimensional spaces and then used to enforce intermediate supervision on the segmentation component. While, the image segmentation component (176.6 MB) is expected to have the capability of capturing fine structures and\nhandle large size e.g., 1024 × 1024, inputs with affordable memory and time costs. In the following experiment, we\nchoose ***U2-Net*** as the image segmentation component because of its strong capability in capturing fine structures.\n\nNote that other segmentation models, such as transformer backbone, are also compatible with this strategy.","metadata":{}},{"cell_type":"code","source":"def build_model(hypar,device):\n    net = hypar[\"model\"]#GOSNETINC(3,1)\n\n    # convert to half precision\n    if(hypar[\"model_digit\"]==\"half\"):\n        net.half()\n        for layer in net.modules():\n            if isinstance(layer, nn.BatchNorm2d):\n                layer.float()\n\n    net.to(device)\n\n    if(hypar[\"restore_model\"]!=\"\"):\n        net.load_state_dict(torch.load(hypar[\"restore_model\"],map_location=device))\n        net.to(device)\n        print(\"model restored\")\n    else:\n        print(\"training from scratch\")\n#     net.eval()  \n    return net","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:40:02.546376Z","iopub.execute_input":"2022-09-16T03:40:02.546803Z","iopub.status.idle":"2022-09-16T03:40:02.556350Z","shell.execute_reply.started":"2022-09-16T03:40:02.546765Z","shell.execute_reply":"2022-09-16T03:40:02.555028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = build_model(hypar,device)","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:41:52.054939Z","iopub.execute_input":"2022-09-16T03:41:52.055335Z","iopub.status.idle":"2022-09-16T03:41:53.837903Z","shell.execute_reply.started":"2022-09-16T03:41:52.055301Z","shell.execute_reply":"2022-09-16T03:41:53.836831Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Optimizer**","metadata":{}},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=25e-4, betas=(0.9, 0.999), eps=1e-08, weight_decay=0)","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:42:00.538088Z","iopub.execute_input":"2022-09-16T03:42:00.538837Z","iopub.status.idle":"2022-09-16T03:42:00.546245Z","shell.execute_reply.started":"2022-09-16T03:42:00.538800Z","shell.execute_reply":"2022-09-16T03:42:00.545292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Scheduler**","metadata":{}},{"cell_type":"markdown","source":"## **Evaluation fucntions**","metadata":{}},{"cell_type":"code","source":"def mae_torch(pred,gt):\n\n\th,w = gt.shape[0:2]\n\tsumError = torch.sum(torch.absolute(torch.sub(pred.float(), gt.float())))\n\tmaeError = torch.divide(sumError,float(h)*float(w)*255.0+1e-4)\n\n\treturn maeError\n\n\ndef f1score_torch(pd,gt):\n\n\t# print(gt.shape)\n\tgtNum = torch.sum((gt>128).float()*1) ## number of ground truth pixels\n\n\tpp = pd[gt>128]\n\tnn = pd[gt<=128]\n\n\tpp_hist =torch.histc(pp,bins=255,min=0,max=255)\n\tnn_hist = torch.histc(nn,bins=255,min=0,max=255)\n\n\n\tpp_hist_flip = torch.flipud(pp_hist)\n\tnn_hist_flip = torch.flipud(nn_hist)\n\n\tpp_hist_flip_cum = torch.cumsum(pp_hist_flip, dim=0)\n\tnn_hist_flip_cum = torch.cumsum(nn_hist_flip, dim=0)\n\n\tprecision = (pp_hist_flip_cum)/(pp_hist_flip_cum + nn_hist_flip_cum + 1e-4)#torch.divide(pp_hist_flip_cum,torch.sum(torch.sum(pp_hist_flip_cum, nn_hist_flip_cum), 1e-4))\n\trecall = (pp_hist_flip_cum)/(gtNum + 1e-4)\n\tf1 = (1+0.3)*precision*recall/(0.3*precision+recall + 1e-4)\n\n\treturn torch.reshape(precision,(1,precision.shape[0])),torch.reshape(recall,(1,recall.shape[0])),torch.reshape(f1,(1,f1.shape[0]))\n\n\n\n\ndef f1_mae_torch(pred, gt, valid_dataset, idx, mybins, hypar):\n\n\timport time\n\ttic = time.time()\n\n\tif(len(gt.shape)>2):\n\t\tgt = gt[:,:,0]\n\n\tpre, rec, f1 = f1score_torch(pred,gt)\n\tmae = mae_torch(pred,gt)\n\n\n\t# hypar[\"valid_out_dir\"] = hypar[\"valid_out_dir\"]+\"-eval\" ###\n\tif(hypar[\"valid_out_dir\"]!=\"\"):\n\t\tif(not os.path.exists(hypar[\"valid_out_dir\"])):\n\t\t\tos.mkdir(hypar[\"valid_out_dir\"])\n\t\tdataset_folder = os.path.join(hypar[\"valid_out_dir\"],valid_dataset.dataset[\"data_name\"][idx])\n\t\tif(not os.path.exists(dataset_folder)):\n\t\t\tos.mkdir(dataset_folder)\n\t\tio.imsave(os.path.join(dataset_folder,valid_dataset.dataset[\"im_name\"][idx]+\".png\"),pred.cpu().data.numpy().astype(np.uint8))\n# \tprint(valid_dataset.dataset[\"im_name\"][idx]+\".png\")\n# \tprint(\"time for evaluation : \", time.time()-tic)\n\n\treturn pre.cpu().data.numpy(), rec.cpu().data.numpy(), f1.cpu().data.numpy(), mae.cpu().data.numpy()","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:42:01.977333Z","iopub.execute_input":"2022-09-16T03:42:01.978277Z","iopub.status.idle":"2022-09-16T03:42:01.992643Z","shell.execute_reply.started":"2022-09-16T03:42:01.978238Z","shell.execute_reply":"2022-09-16T03:42:01.991477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Training Function**","metadata":{}},{"cell_type":"markdown","source":"### **Intermediate Supervision Network training fucntion**","metadata":{}},{"cell_type":"code","source":"\ndef get_gt_encoder(train_dataloaders, train_datasets, valid_dataloaders, valid_datasets, \n                   hypar, train_dataloaders_val, train_datasets_val): #model_path, model_save_fre, max_ite=1000000):\n\n\n    torch.manual_seed(hypar[\"seed\"])\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(hypar[\"seed\"])\n\n    print(\"define gt encoder ...\")\n    net = ISNetGTEncoder() #UNETGTENCODERCombine()\n    ## load the existing model gt encoder\n    if(hypar[\"gt_encoder_model\"]!=\"\"):\n        model_path = hypar[\"gt_encoder_model\"]\n        if torch.cuda.is_available():\n            net.load_state_dict(torch.load(model_path))\n            net.cuda()\n            print(\"gt encoder restored from the saved weights ...\")\n#             print(\"gt_model_restored\")\n        else:\n            net.load_state_dict(torch.load(model_path,map_location=\"cpu\"))\n            print(\"gt encoder restored from the saved weights ...\")  \n#         return net ############\n\n    if torch.cuda.is_available():\n        net.cuda()\n\n    print(\"--- define optimizer for GT Encoder---\")\n    optimizer = optim.Adam(net.parameters(), lr=1e-3, betas=(0.9, 0.999), eps=1e-08, weight_decay=0)\n#     scheduler_gt = lr_scheduler.CosineAnnealingLR(optimizer,T_max=2032, eta_min=1e-5)\n    scheduler_gt = lr_scheduler.CosineAnnealingWarmRestarts(optimizer, 271, T_mult=2, eta_min=6e-6)\n    \n    \n\n    model_path = hypar[\"model_path\"]\n    model_save_fre = hypar[\"model_save_fre\"]\n    max_ite = hypar[\"max_ite\"]\n    batch_size_train = hypar[\"batch_size_train\"]\n    batch_size_valid = hypar[\"batch_size_valid\"]\n\n    if(not os.path.exists(model_path)):\n        os.mkdir(model_path)\n\n    ite_num = hypar[\"start_ite\"] # count the total iteration number\n    ite_num4val = 0 #\n    running_loss = 0.0 # count the toal loss\n    running_tar_loss = 0.0 # count the target output loss\n    last_f1 = [0 for x in range(len(valid_dataloaders))]\n\n    train_num = train_datasets[0].__len__()\n\n    net.train()\n\n    start_last = time.time()\n    gos_dataloader = train_dataloaders[0]\n    epoch_num = hypar[\"max_epoch_num_gte\"]\n    notgood_cnt = 0\n    for epoch in range(epoch_num): ## set the epoch num as 100000\n\n        for i, data in enumerate(gos_dataloader):\n\n            if(ite_num >= max_ite):\n                print(\"Training Reached the Maximal Iteration Number \", max_ite)\n                exit()\n\n            # start_read = time.time()\n            ite_num = ite_num + 1\n            ite_num4val = ite_num4val + 1\n\n            # get the inputs\n            labels = data['label'] #*255.0\n\n            if(hypar[\"model_digit\"]==\"full\"):\n                labels = labels.type(torch.FloatTensor)\n            else:\n                labels = labels.type(torch.HalfTensor)\n\n            # wrap them in Variable\n            if torch.cuda.is_available():\n                labels_v = Variable(labels.cuda(), requires_grad=False)\n            else:\n                labels_v = Variable(labels, requires_grad=False)\n\n            # print(\"time lapse for data preparation: \", time.time()-start_read, ' s')\n\n            # y zero the parameter gradients\n            start_inf_loss_back = time.time()\n            optimizer.zero_grad()\n\n            ds, fs = net(labels_v)#net(inputs_v)\n            loss2, loss = net.compute_loss(ds, labels_v)\n\n            loss.backward()\n            optimizer.step()\n            scheduler_gt.step()\n#             if scheduler is not None:\n\n            running_loss += loss.item()\n            running_tar_loss += loss2.item()\n\n            # del outputs, loss\n            del ds, loss2, loss\n            end_inf_loss_back = time.time()-start_inf_loss_back\n\n#             if ite_num % 2167 == 0:\n        current_lr = optimizer.param_groups[0]['lr']\n        print(\"GT Encoder Training>>>\"+model_path.split('/')[-1]+\" - [epoch: %3d/%3d, batch: %5d/%5d, ite: %d] train loss: %3f, tar: %3f, time-per-iter: %3f s, time_read: %3f, lr: %6f\" % (\n        epoch + 1, epoch_num, (i + 1) * batch_size_train, train_num, ite_num, running_loss / ite_num4val, running_tar_loss / ite_num4val, time.time()-start_last, time.time()-start_last-end_inf_loss_back, current_lr))\n        start_last = time.time()\n\n        if (epoch+1) % 50 == 0:  # validate every 2000 iterations\n            notgood_cnt += 1\n            # net.eval()\n            # tmp_f1, tmp_mae, val_loss, tar_loss, i_val, tmp_time = valid_gt_encoder(net, valid_dataloaders, valid_datasets, hypar, epoch)\n            tmp_f1, tmp_mae, val_loss, tar_loss, i_val, tmp_time = valid_gt_encoder(net, train_dataloaders_val, train_datasets_val, hypar, epoch)\n\n            net.train()  # resume train\n\n            tmp_out = 0\n            print(\"last_f1:\",last_f1)\n            print(\"tmp_f1:\",tmp_f1)\n            for fi in range(len(last_f1)):\n                if(tmp_f1[fi]>last_f1[fi]):\n                    tmp_out = 1\n            print(\"tmp_out:\",tmp_out)\n            if(tmp_out):\n                notgood_cnt = 0\n                last_f1 = tmp_f1\n                tmp_f1_str = [str(round(f1x,4)) for f1x in tmp_f1]\n                tmp_mae_str = [str(round(mx,4)) for mx in tmp_mae]\n                maxf1 = '_'.join(tmp_f1_str)\n                meanM = '_'.join(tmp_mae_str)\n                # .cpu().detach().numpy()\n                model_name = \"/GTENCODER\"+str(ite_num)+\\\n                            \"_traLoss_\"+str(np.round(running_loss / ite_num4val,4))+\\\n                            \"_valLoss_\"+str(np.round(val_loss /(i_val+1),4))+\\\n                            \"_maxF1_\" + maxf1 + \\\n                            \"_mae_\" + meanM + \".pth\"\n                torch.save(net.state_dict(), model_path + model_name)\n                print(\"gt_model_saved\")\n\n            running_loss = 0.0\n            running_tar_loss = 0.0\n            ite_num4val = 0\n\n            if(tmp_f1[0]>0.99):\n                print(\"GT encoder is well-trained and obtained...\")\n                return net\n\n            if(notgood_cnt >= hypar[\"early_stop\"]):\n                print(\"No improvements in the last \"+str(notgood_cnt)+\" validation periods, so training stopped !\")\n                exit()\n\n    print(\"Training Reaches The Maximum Epoch Number\")\n    return net","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:44:30.393468Z","iopub.execute_input":"2022-09-16T03:44:30.394063Z","iopub.status.idle":"2022-09-16T03:44:30.418465Z","shell.execute_reply.started":"2022-09-16T03:44:30.394025Z","shell.execute_reply":"2022-09-16T03:44:30.417327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef valid_gt_encoder(net, valid_dataloaders, valid_datasets, hypar, epoch=0):\n    net.eval()\n    print(\"Validating...\")\n    epoch_num = hypar[\"max_epoch_num_gte\"]\n\n    val_loss = 0.0\n    tar_loss = 0.0\n\n\n    tmp_f1 = []\n    tmp_mae = []\n    tmp_time = []\n\n    start_valid = time.time()\n    for k in range(len(valid_dataloaders)):\n\n        valid_dataloader = valid_dataloaders[k]\n        valid_dataset = valid_datasets[k]\n\n        val_num = valid_dataset.__len__()\n        mybins = np.arange(0,256)\n        PRE = np.zeros((val_num,len(mybins)-1))\n        REC = np.zeros((val_num,len(mybins)-1))\n        F1 = np.zeros((val_num,len(mybins)-1))\n        MAE = np.zeros((val_num))\n\n        val_cnt = 0.0\n        i_val = None\n\n        for i_val, data_val in enumerate(valid_dataloader):\n\n#             imidx_val, inputs_val, labels_val, shapes_val = data_val['imidx'], data_val['image'], data_val['label'], data_val['shape']\n#             imidx_val, labels_val, shapes_val = data_val['imidx'], data_val['label']*255.0, data_val['shape']\n            imidx_val, labels_val, shapes_val = data_val['imidx'], data_val['label'], data_val['shape']\n            \n            \n            if(hypar[\"model_digit\"]==\"full\"):\n                labels_val = labels_val.type(torch.FloatTensor)\n            else:\n                labels_val = labels_val.type(torch.HalfTensor)\n\n            # wrap them in Variable\n            if torch.cuda.is_available():\n                labels_val_v = Variable(labels_val.cuda(), requires_grad=False)\n            else:\n                labels_val_v = Variable(labels_val,requires_grad=False)\n\n            t_start = time.time()\n            ds_val = net(labels_val_v)[0]\n            t_end = time.time()-t_start\n            tmp_time.append(t_end)\n\n            # loss2_val, loss_val = muti_loss_fusion(ds_val, labels_val_v)\n            loss2_val, loss_val = net.compute_loss(ds_val, labels_val_v)\n\n            # compute F measure\n            for t in range(hypar[\"batch_size_valid\"]):\n                val_cnt = val_cnt + 1.0\n#                 print(\"num of val: \", val_cnt)\n                i_test = imidx_val[t].data.numpy()\n\n                pred_val = ds_val[0][t,:,:,:] # B x 1 x H x W\n\n                ## recover the prediction spatial size to the orignal image size\n                pred_val = torch.squeeze(F.upsample(torch.unsqueeze(pred_val,0),(shapes_val[t][0],shapes_val[t][1]),mode='bilinear'))\n\n                ma = torch.max(pred_val)\n                mi = torch.min(pred_val)\n                pred_val = (pred_val-mi)/(ma-mi) # max = 1\n                # pred_val = normPRED(pred_val)\n\n                gt = np.squeeze(io.imread(valid_dataset.dataset[\"ori_gt_path\"][i_test])) # max = 255\n                with torch.no_grad():\n                    gt = torch.tensor(gt).to(device)\n\n#                 pre,rec,f1,mae = f1_mae_torch(pred_val*255, gt, valid_dataset, i_test, mybins, hypar)\n                pre,rec,f1,mae = f1_mae_torch(pred_val*255, gt*255.0, valid_dataset, i_test, mybins, hypar)\n                \n                \n                PRE[i_test,:]=pre\n                REC[i_test,:] = rec\n                F1[i_test,:] = f1\n                MAE[i_test] = mae\n\n            del ds_val, gt\n            gc.collect()\n            torch.cuda.empty_cache()\n\n            # if(loss_val.data[0]>1):\n            val_loss += loss_val.item()#data[0]\n            tar_loss += loss2_val.item()#data[0]\n            \n            del loss2_val, loss_val\n            \n#             if i_val % 700 == 0:\n\n        print('============================')\n        print(\"[validating: %5d/%5d] val_ls:%f, tar_ls: %f, f1: %f, mae: %f, time: %f\"% (i_val, val_num, val_loss / (i_val + 1), tar_loss / (i_val + 1), np.amax(F1[i_test,:]), MAE[i_test],t_end))\n    \n        PRE_m = np.mean(PRE,0)\n        REC_m = np.mean(REC,0)\n        f1_m = (1+0.3)*PRE_m*REC_m/(0.3*PRE_m+REC_m+1e-8)\n        # print('--------------:', np.mean(f1_m))\n        tmp_f1.append(np.amax(f1_m))\n        tmp_mae.append(np.mean(MAE))\n        print(\"The max F1 Score: %f\"%(np.max(f1_m)))\n        print(\"MAE: \", np.mean(MAE))\n\n    # print('[epoch: %3d/%3d, ite: %5d] tra_ls: %3f, val_ls: %3f, tar_ls: %3f, maxf1: %3f, val_time: %6f'% (epoch + 1, epoch_num, ite_num, running_loss / ite_num4val, val_loss/val_cnt, tar_loss/val_cnt, tmp_f1[-1], time.time()-start_valid))\n\n    return tmp_f1, tmp_mae, val_loss, tar_loss, i_val, tmp_time","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:44:33.738241Z","iopub.execute_input":"2022-09-16T03:44:33.738662Z","iopub.status.idle":"2022-09-16T03:44:33.757301Z","shell.execute_reply.started":"2022-09-16T03:44:33.738630Z","shell.execute_reply":"2022-09-16T03:44:33.756283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Main model training**","metadata":{}},{"cell_type":"code","source":"def train(net, optimizer, train_dataloaders, train_datasets, valid_dataloaders, valid_datasets, \n          hypar, train_dataloaders_val, train_datasets_val): #model_path, model_save_fre, max_ite=1000000):\n\n    if hypar[\"interm_sup\"]:\n        print(\"Get the gt encoder ...\")\n        featurenet = get_gt_encoder(train_dataloaders, train_datasets, valid_dataloaders, \n                                    valid_datasets, hypar,train_dataloaders_val, train_datasets_val)\n        ## freeze the weights of gt encoder\n        for param in featurenet.parameters():\n            param.requires_grad=False\n\n#     scheduler = lr_scheduler.CosineAnnealingLR(optimizer,T_max=3523, eta_min=6e-6)\n    scheduler = lr_scheduler.CosineAnnealingWarmRestarts(optimizer, 542, T_mult=2, eta_min=6e-6)\n    model_path = hypar[\"model_path\"]\n    model_save_fre = hypar[\"model_save_fre\"]\n    max_ite = hypar[\"max_ite\"]\n    batch_size_train = hypar[\"batch_size_train\"]\n    batch_size_valid = hypar[\"batch_size_valid\"]\n\n    if(not os.path.exists(model_path)):\n        os.mkdir(model_path)\n\n    ite_num = hypar[\"start_ite\"] # count the toal iteration number\n    ite_num4val = 0 #\n    running_loss = 0.0 # count the toal loss\n    running_tar_loss = 0.0 # count the target output loss\n    last_f1 = [0 for x in range(len(valid_dataloaders))]\n\n    train_num = train_datasets[0].__len__()\n\n    net.train()\n\n    start_last = time.time()\n    gos_dataloader = train_dataloaders[0]\n    iters = len(gos_dataloader)\n    epoch_num = hypar[\"max_epoch_num\"]\n    notgood_cnt = 0\n    for epoch in range(epoch_num): ## set the epoch num as 100000\n\n        for i, data in enumerate(gos_dataloader):\n\n            # start_read = time.time()\n            ite_num = ite_num + 1\n            ite_num4val = ite_num4val + 1\n\n            # get the inputs\n#             inputs, labels = data['image'], data['label']*255.0\n            inputs, labels = data['image'], data['label']  \n\n\n            if(hypar[\"model_digit\"]==\"full\"):\n                inputs = inputs.type(torch.FloatTensor)\n                labels = labels.type(torch.FloatTensor)\n            else:\n                inputs = inputs.type(torch.HalfTensor)\n                labels = labels.type(torch.HalfTensor)\n\n            # wrap them in Variable\n            if torch.cuda.is_available():\n                inputs_v, labels_v = Variable(inputs.cuda(), requires_grad=False), Variable(labels.cuda(), requires_grad=False)\n            else:\n                inputs_v, labels_v = Variable(inputs, requires_grad=False), Variable(labels, requires_grad=False)\n\n            # print(\"time lapse for data preparation: \", time.time()-start_read, ' s')\n\n            # y zero the parameter gradients\n            start_inf_loss_back = time.time()\n            optimizer.zero_grad()\n\n            if hypar[\"interm_sup\"]:\n                # forward + backward + optimize\n                ds,dfs = net(inputs_v)\n                _,fs = featurenet(labels_v) ## extract the gt encodings\n                loss2, loss = net.compute_loss_kl(ds, labels_v, dfs, fs, mode='MSE')\n            else:\n                # forward + backward + optimize\n                ds,_ = net(inputs_v)\n                loss2, loss = net.compute_loss(ds, labels_v)\n\n            loss.backward()\n            optimizer.step()\n            scheduler.step(epoch + i / iters)            \n#             if scheduler is not None:\n\n            # # print statistics\n            running_loss += loss.item()\n            running_tar_loss += loss2.item()\n\n            # del outputs, loss\n            del ds, loss2, loss\n            end_inf_loss_back = time.time()-start_inf_loss_back\n            \n#             if ite_num % 2167 == 0:\n        current_lr = optimizer.param_groups[0]['lr']\n        print(\">>>\"+model_path.split('/')[-1]+\" - [epoch: %3d/%3d, batch: %5d/%5d, ite: %d] train loss: %3f, tar: %3f, time-per-epoch: %3f s, time_read: %3f, lr: %6f\" % (\n        epoch + 1, epoch_num, (i + 1) * batch_size_train, train_num, ite_num, running_loss / ite_num4val, running_tar_loss / ite_num4val, time.time()-start_last, time.time()-start_last-end_inf_loss_back, current_lr))\n            \n        start_last = time.time()\n\n        if (epoch +1) % 10 == 0:  # validate every 5 epochs\n            notgood_cnt += 1\n            net.eval()\n            tmp_f1, tmp_mae, val_loss, tar_loss, i_val, tmp_time = valid(net, valid_dataloaders, valid_datasets, hypar, epoch)\n            net.train()  # resume train\n\n            tmp_out = 0\n            print(\"last_f1:\",last_f1)\n            print(\"tmp_f1:\",tmp_f1)\n            for fi in range(len(last_f1)):\n                if(tmp_f1[fi]>last_f1[fi]):\n                    tmp_out = 1\n            print(\"tmp_out:\",tmp_out)\n            if(tmp_out):\n                notgood_cnt = 0\n                last_f1 = tmp_f1\n                tmp_f1_str = [str(round(f1x,4)) for f1x in tmp_f1]\n                tmp_mae_str = [str(round(mx,4)) for mx in tmp_mae]\n                maxf1 = '_'.join(tmp_f1_str)\n                meanM = '_'.join(tmp_mae_str)\n                # .cpu().detach().numpy()\n                model_name = \"/ISNet\"+\\\n                            \"_traLoss_\"+str(np.round(running_loss / ite_num4val,4))+\\\n                            \"_valLoss_\"+str(np.round(val_loss /(i_val+1),4))+\\\n                            \"_maxF1_\" + maxf1 + \\\n                            \"_mae_\" + meanM + \".pth\"\n                torch.save(net.state_dict(), model_path + model_name)\n                print(\"model saved!\")\n\n            running_loss = 0.0\n            running_tar_loss = 0.0\n            ite_num4val = 0\n\n            if(notgood_cnt >= hypar[\"early_stop\"]):\n                print(\"No improvements in the last \"+str(notgood_cnt)+\" validation periods, so training stopped !\")\n                exit()\n\n    print(\"Training Reaches The Maximum Epoch Number\")\n","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:42:05.044038Z","iopub.execute_input":"2022-09-16T03:42:05.044433Z","iopub.status.idle":"2022-09-16T03:42:05.069163Z","shell.execute_reply.started":"2022-09-16T03:42:05.044399Z","shell.execute_reply":"2022-09-16T03:42:05.068035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### **Validation Function**","metadata":{}},{"cell_type":"code","source":"def valid(net, valid_dataloaders, valid_datasets, hypar, epoch=0):\n    net.eval()\n    print(\"Validating...\")\n    epoch_num = hypar[\"max_epoch_num\"]\n\n    val_loss = 0.0\n    tar_loss = 0.0\n    val_cnt = 0.0\n\n    tmp_f1 = []\n    tmp_mae = []\n    tmp_time = []\n\n    start_valid = time.time()\n\n    for k in range(len(valid_dataloaders)):\n\n        valid_dataloader = valid_dataloaders[k]\n        valid_dataset = valid_datasets[k]\n\n        val_num = valid_dataset.__len__()\n        mybins = np.arange(0,256)\n        PRE = np.zeros((val_num,len(mybins)-1))\n        REC = np.zeros((val_num,len(mybins)-1))\n        F1 = np.zeros((val_num,len(mybins)-1))\n        MAE = np.zeros((val_num))\n\n        for i_val, data_val in enumerate(valid_dataloader):\n            val_cnt = val_cnt + 1.0\n            imidx_val, inputs_val, labels_val, shapes_val = data_val['imidx'], data_val['image'], data_val['label'], data_val['shape']\n#             imidx_val, inputs_val, labels_val, shapes_val = data_val['imidx'], data_val['image'], data_val['label']*255.0, data_val['shape']\n            \n\n            if(hypar[\"model_digit\"]==\"full\"):\n                inputs_val = inputs_val.type(torch.FloatTensor)\n                labels_val = labels_val.type(torch.FloatTensor)\n            else:\n                inputs_val = inputs_val.type(torch.HalfTensor)\n                labels_val = labels_val.type(torch.HalfTensor)\n\n            # wrap them in Variable\n            if torch.cuda.is_available():\n                inputs_val_v, labels_val_v = Variable(inputs_val.cuda(), requires_grad=False), Variable(labels_val.cuda(), requires_grad=False)\n            else:\n                inputs_val_v, labels_val_v = Variable(inputs_val, requires_grad=False), Variable(labels_val,requires_grad=False)\n\n            t_start = time.time()\n            ds_val = net(inputs_val_v)[0]\n            t_end = time.time()-t_start\n            tmp_time.append(t_end)\n\n            # loss2_val, loss_val = muti_loss_fusion(ds_val, labels_val_v)\n            loss2_val, loss_val = net.compute_loss(ds_val, labels_val_v)\n\n            # compute F measure\n            for t in range(hypar[\"batch_size_valid\"]):\n                i_test = imidx_val[t].data.numpy()\n\n                pred_val = ds_val[0][t,:,:,:] # B x 1 x H x W\n\n                ## recover the prediction spatial size to the orignal image size\n                pred_val = torch.squeeze(F.upsample(torch.unsqueeze(pred_val,0),(shapes_val[t][0],shapes_val[t][1]),mode='bilinear'))\n\n                # pred_val = normPRED(pred_val)\n                ma = torch.max(pred_val)\n                mi = torch.min(pred_val)\n                pred_val = (pred_val-mi)/(ma-mi) # max = 1\n\n                if len(valid_dataset.dataset[\"ori_gt_path\"]) != 0:\n                    gt = np.squeeze(io.imread(valid_dataset.dataset[\"ori_gt_path\"][i_test])) # max = 255\n                else:\n                    gt = np.zeros((shapes_val[t][0],shapes_val[t][1]))\n                with torch.no_grad():\n                    gt = torch.tensor(gt).to(device)\n\n#                 pre,rec,f1,mae = f1_mae_torch(pred_val*255, gt, valid_dataset, i_test, mybins, hypar)\n                pre,rec,f1,mae = f1_mae_torch(pred_val*255, gt*255, valid_dataset, i_test, mybins, hypar)\n\n\n                PRE[i_test,:]=pre\n                REC[i_test,:] = rec\n                F1[i_test,:] = f1\n                MAE[i_test] = mae\n\n                del ds_val, gt\n                gc.collect()\n                torch.cuda.empty_cache()\n\n            # if(loss_val.data[0]>1):\n            val_loss += loss_val.item()#data[0]\n            tar_loss += loss2_val.item()#data[0]\n            \n            del loss2_val, loss_val\n            \n#             if i_val % 350 == 0:\n\n        print('============================')\n        print(\"[validating: %5d/%5d] val_ls:%f, tar_ls: %f, f1: %f, mae: %f, time: %f\"% (i_val, val_num, val_loss / (i_val + 1), tar_loss / (i_val + 1), np.amax(F1[i_test,:]), MAE[i_test],t_end))\n        PRE_m = np.mean(PRE,0)\n        REC_m = np.mean(REC,0)\n        f1_m = (1+0.3)*PRE_m*REC_m/(0.3*PRE_m+REC_m+1e-8)\n\n        tmp_f1.append(np.amax(f1_m))\n        tmp_mae.append(np.mean(MAE))\n\n    return tmp_f1, tmp_mae, val_loss, tar_loss, i_val, tmp_time","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:42:30.215306Z","iopub.execute_input":"2022-09-16T03:42:30.215865Z","iopub.status.idle":"2022-09-16T03:42:30.250791Z","shell.execute_reply.started":"2022-09-16T03:42:30.215819Z","shell.execute_reply":"2022-09-16T03:42:30.249852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_dataloader = train_dataloaders[0]\n# train_dataset = train_datasets[0]\n# val_num = train_dataset.__len__()\n# mybins = np.arange(0,256)\n# dataitr = iter(train_dataloader)","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:42:31.283600Z","iopub.execute_input":"2022-09-16T03:42:31.283962Z","iopub.status.idle":"2022-09-16T03:42:31.288193Z","shell.execute_reply.started":"2022-09-16T03:42:31.283930Z","shell.execute_reply":"2022-09-16T03:42:31.287226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Run Training**","metadata":{}},{"cell_type":"code","source":"train(model,\n      optimizer,\n      train_dataloaders,\n      train_datasets,\n      valid_dataloaders,\n      valid_datasets,\n      hypar,\n      train_dataloaders_val, train_datasets_val)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-09-16T03:45:21.176835Z","iopub.execute_input":"2022-09-16T03:45:21.177202Z","iopub.status.idle":"2022-09-16T05:29:00.604948Z","shell.execute_reply.started":"2022-09-16T03:45:21.177170Z","shell.execute_reply":"2022-09-16T05:29:00.603788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), \"./ISNet_latest.pth\")\nprint(\"model saved!\")","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:40:08.617465Z","iopub.status.idle":"2022-09-16T03:40:08.618844Z","shell.execute_reply.started":"2022-09-16T03:40:08.618567Z","shell.execute_reply":"2022-09-16T03:40:08.618598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls ./saved_models/","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:40:08.621543Z","iopub.status.idle":"2022-09-16T03:40:08.623972Z","shell.execute_reply.started":"2022-09-16T03:40:08.623672Z","shell.execute_reply":"2022-09-16T03:40:08.623700Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# `!mv ./saved_models/GTENCODER-gpu_itr_18000_traLoss_0.0536_traTarLoss_0.0001_valLoss_0.0528_valTarLoss_0.0001_maxF1_0.2134_mae_0.1074_time_0.024927.pth ./gt_encoder.pth","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:40:08.626331Z","iopub.status.idle":"2022-09-16T03:40:08.627934Z","shell.execute_reply.started":"2022-09-16T03:40:08.627630Z","shell.execute_reply":"2022-09-16T03:40:08.627656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !mv ./saved_models/GTENCODER-gpu_itr_4065_traLoss_0.0639_traTarLoss_0.0004_valLoss_0.058_valTarLoss_0.0002_maxF1_0.2083_mae_0.1087_time_0.024932.pth ./gt_encoder.pth","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:40:08.629684Z","iopub.status.idle":"2022-09-16T03:40:08.634679Z","shell.execute_reply.started":"2022-09-16T03:40:08.630263Z","shell.execute_reply":"2022-09-16T03:40:08.630291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls ","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:40:08.636132Z","iopub.status.idle":"2022-09-16T03:40:08.636976Z","shell.execute_reply.started":"2022-09-16T03:40:08.636701Z","shell.execute_reply":"2022-09-16T03:40:08.636727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !rm ./saved_models/GTENCODER-gpu_itr_6775_traLoss_0.0684_traTarLoss_0.0004_valLoss_0.0602_valTarLoss_0.0002_maxF1_0.2103_mae_0.1082_time_0.024483.pth #remove previous checkpoints after moving the final checkpoints to another dir","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:40:08.638298Z","iopub.status.idle":"2022-09-16T03:40:08.643553Z","shell.execute_reply.started":"2022-09-16T03:40:08.643277Z","shell.execute_reply":"2022-09-16T03:40:08.643304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -r cache_dir cache_dir_val","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:40:08.644999Z","iopub.status.idle":"2022-09-16T03:40:08.651321Z","shell.execute_reply.started":"2022-09-16T03:40:08.647738Z","shell.execute_reply":"2022-09-16T03:40:08.647764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !ls","metadata":{"execution":{"iopub.status.busy":"2022-09-16T03:40:08.652772Z","iopub.status.idle":"2022-09-16T03:40:08.653726Z","shell.execute_reply.started":"2022-09-16T03:40:08.653314Z","shell.execute_reply":"2022-09-16T03:40:08.653341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}