{"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":"#### This notebook is heavily (almost a rip off) based on https://www.kaggle.com/code/thedevastator/training-fastai-baseline -- thanks to https://www.kaggle.com/thedevastator\n\n\n##### Since the images are from healhty people, I wanted to see if we can determine Age/Gender using only the biopsy images.","metadata":{}},{"cell_type":"code","source":"%reload_ext autoreload\n%autoreload 2\n%matplotlib inline\n\nfrom fastai.vision.all import *\nfrom torch.utils.data import Dataset, DataLoader\nimport pandas as pd\nimport numpy as np\nimport os\nimport cv2\nimport gc\nimport random\nfrom albumentations import *\nfrom sklearn.model_selection import KFold\nimport matplotlib.pyplot as plt\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":3.556569,"end_time":"2021-03-11T18:12:54.195121","exception":false,"start_time":"2021-03-11T18:12:50.638552","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-24T10:26:42.055694Z","iopub.execute_input":"2022-06-24T10:26:42.056145Z","iopub.status.idle":"2022-06-24T10:26:42.228989Z","shell.execute_reply.started":"2022-06-24T10:26:42.056106Z","shell.execute_reply":"2022-06-24T10:26:42.227980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bs = 64\nnfolds = 4\nfold = 0\nSEED = 2020\nroot_data = '../input/hubmap-2022-512x512/'\nTRAIN = f'{root_data}/train/'\nMASKS = f'{root_data}/masks/'\nLABELS = '../input/hubmap-organ-segmentation/train.csv'\nNUM_WORKERS = 4\ndevice = torch.device(f\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"papermill":{"duration":0.045568,"end_time":"2021-03-11T18:12:54.252135","exception":false,"start_time":"2021-03-11T18:12:54.206567","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-24T10:26:42.234428Z","iopub.execute_input":"2022-06-24T10:26:42.236794Z","iopub.status.idle":"2022-06-24T10:26:42.361152Z","shell.execute_reply.started":"2022-06-24T10:26:42.236754Z","shell.execute_reply":"2022-06-24T10:26:42.359914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    #the following line gives ~10% speedup\n    #but may lead to some stochasticity in the results \n    torch.backends.cudnn.benchmark = True\n    return seed\n    \nseed_everything(SEED)","metadata":{"_kg_hide-input":true,"papermill":{"duration":0.04608,"end_time":"2021-03-11T18:12:54.308988","exception":false,"start_time":"2021-03-11T18:12:54.262908","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-24T10:26:42.367758Z","iopub.execute_input":"2022-06-24T10:26:42.368384Z","iopub.status.idle":"2022-06-24T10:26:42.503177Z","shell.execute_reply.started":"2022-06-24T10:26:42.368339Z","shell.execute_reply":"2022-06-24T10:26:42.501979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://www.kaggle.com/datasets/thedevastator/hubmap-2022-256x256\nmean = np.array([0.7720342, 0.74582646, 0.76392896])\nstd = np.array([0.24745085, 0.26182273, 0.25782376])\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, is_debug=False):\n        ids = pd.read_csv(LABELS).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        self.df = pd.read_csv(LABELS)\n        self.is_debug = is_debug\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        X = img2tensor((img/255.0 - mean)/std)\n        Xmasks = img2tensor(mask)\n        \n        _id = int(os.path.splitext(fname)[0].split('_')[0])\n        row = self.df[self.df['id'] == _id].reset_index(drop=True)\n        y = torch.tensor([row['age'].values[0], 0 if row['sex'].values[0] == 'Male' else 1], dtype=torch.float32)\n        \n        if self.is_debug:\n            return X, Xmasks, y\n        return X, y\n    \ndef get_aug(p=1.0):\n    return Compose([\n        HorizontalFlip(),\n        VerticalFlip(),\n        RandomRotate90(),\n        ShiftScaleRotate(\n            shift_limit=0.0625,\n            scale_limit=0.2,\n            rotate_limit=15,\n            p=0.5, \n            border_mode=cv2.BORDER_CONSTANT\n        ),\n        OneOf([\n            OpticalDistortion(p=0.3),\n            GridDistortion(p=.1),\n            IAAPiecewiseAffine(p=0.3),\n        ], p=0.), # switched off\n        OneOf([\n            HueSaturationValue(10,15,10),\n            CLAHE(clip_limit=2),\n            RandomBrightnessContrast(),            \n        ], p=0.), # switched off\n    ], p=p)","metadata":{"papermill":{"duration":0.057765,"end_time":"2021-03-11T18:12:54.40978","exception":false,"start_time":"2021-03-11T18:12:54.352015","status":"completed"},"tags":[],"_kg_hide-input":false,"execution":{"iopub.status.busy":"2022-06-24T10:26:42.509123Z","iopub.execute_input":"2022-06-24T10:26:42.509572Z","iopub.status.idle":"2022-06-24T10:26:42.655655Z","shell.execute_reply.started":"2022-06-24T10:26:42.509531Z","shell.execute_reply":"2022-06-24T10:26:42.654158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#example of train images with masks\nds = HuBMAPDataset(tfms=get_aug(), is_debug=True)\ndl = DataLoader(ds,batch_size=4,shuffle=False,num_workers=NUM_WORKERS)\nimgs,masks,tgts = 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.title(tgts[i].numpy().tolist())\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":{"papermill":{"duration":12.580982,"end_time":"2021-03-11T18:13:07.001785","exception":false,"start_time":"2021-03-11T18:12:54.420803","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-24T10:26:42.662678Z","iopub.execute_input":"2022-06-24T10:26:42.663127Z","iopub.status.idle":"2022-06-24T10:26:45.278218Z","shell.execute_reply.started":"2022-06-24T10:26:42.663085Z","shell.execute_reply":"2022-06-24T10:26:45.277213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    import timm\nexcept:\n    !pip install timm\n    import timm\n\nclass Model(nn.Module):\n    def __init__(self, model_name='tf_efficientnet_b0_ns', pretrained=True):\n        super().__init__()\n        torch.cuda.empty_cache()\n        self.model = timm.create_model(model_name, pretrained=pretrained, num_classes=2)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x\n\n# model = Model('resnet18', pretrained=True).to(device)","metadata":{"_kg_hide-input":true,"papermill":{"duration":0.110724,"end_time":"2021-03-11T18:13:07.317006","exception":false,"start_time":"2021-03-11T18:13:07.206282","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-24T10:26:45.279434Z","iopub.execute_input":"2022-06-24T10:26:45.279771Z","iopub.status.idle":"2022-06-24T10:27:09.952268Z","shell.execute_reply.started":"2022-06-24T10:26:45.279737Z","shell.execute_reply":"2022-06-24T10:27:09.950882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{"papermill":{"duration":0.040885,"end_time":"2021-03-11T18:13:08.336994","exception":false,"start_time":"2021-03-11T18:13:08.296109","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"### Trying out a resnet18 on 512x512 images [https://www.kaggle.com/datasets/thedevastator/hubmap-2022-512x512]","metadata":{}},{"cell_type":"code","source":"def the_loss(y, y1):\n    gender_loss = F.binary_cross_entropy_with_logits(y[:, 1], y1[:, 1], reduction='mean')\n    age_loss = F.mse_loss(y[:, 0], y1[:, 0], reduction='mean')\n    return age_loss + gender_loss\n\ndef rmse_age(y1, y):\n    y, y1 = y[:, 0], y1[:, 0]\n    return ((y - y1)**2).mean() ** 0.5\n\ndef gender_accuracy(y1, y):\n    y, y1 = y[:, 1], y1[:, 1]\n    y1, y = nn.Sigmoid()(y1.flatten()) < 0.5, y.flatten() == 0\n    return (y==y1).sum()/len(y)\n\nfor fold in range(nfolds):\n    if fold not in [0]: continue\n    ds_t = HuBMAPDataset(fold=fold, train=True, tfms=get_aug())\n    ds_v = HuBMAPDataset(fold=fold, train=False)\n    data = ImageDataLoaders.from_dsets(ds_t,ds_v,bs=bs,num_workers=NUM_WORKERS,pin_memory=True).cuda()\n    model = Model('resnet18', pretrained=True).to(device)\n    learn = Learner(\n        data, model, loss_func=the_loss,\n        metrics=[rmse_age, gender_accuracy],\n        cbs=[\n            ShowGraphCallback(),\n            SaveModelCallback(fname=f'best_fold_{fold}'),\n        ],\n    ).to_fp16()\n\n    learn.fit_one_cycle(10, lr_max=1e-3, pct_start=0.)\n    gc.collect()","metadata":{"_kg_hide-output":false,"papermill":{"duration":29637.81713,"end_time":"2021-03-12T02:27:06.195179","exception":false,"start_time":"2021-03-11T18:13:08.378049","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-06-24T10:27:09.958427Z","iopub.execute_input":"2022-06-24T10:27:09.961246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}