{"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 is a simple implementation of segmentation using CoAT[1] \n\nThe solution is surprisingly simple:\n- there is stricitly no decoder (like upnet, fpn, etc) ... i simple concate the transformer output at different level and fuse with a 3x3 conv2d layer\n- there is no tiling or other complicated post processing \n\n\n[1] Co-Scale Conv-Attentional Image Transformers - Weijian Xu,  ICCV2021 oral  \nhttps://arxiv.org/abs/2104.06399  \n\ncode and credits from  \nhttps://github.com/mlpc-ucsd/CoaT\n\nExpected results :\n\n ![https://i.ibb.co/3s93QwN/Selection-075.png](https://i.ibb.co/3s93QwN/Selection-075.png)\n \n \nI later train for more interations:  \n\n<a href=\"https://ibb.co/c69jWKx\"><img src=\"https://i.ibb.co/HtyjcLX/Selection-064.png\" alt=\"Selection-064\" border=\"0\"></a>\n\n\n<span style=\"color:red\">care needs to be taken if you are using SWA (Stochastic Weight Averaging) to average the weights. There are shared parameters. These should only be averaged once!</span>\n\n","metadata":{}},{"cell_type":"code","source":"import sys, os \nsys.path.append('../input/hubmap-submit-06') \nsys.path.append('../input/hubmap-submit-06/[third_party]')   \n\nfrom kaggle_hubmap_kv3 import *\nimport importlib\nfrom timeit import default_timer as timer\n\nimport torch\nimport torch.cuda.amp as amp\nimport torch.nn.functional as F\nprint('import ok\\n')\n\n#-- configure ---------------------------------------------\nimage_size = 768 #512\n\norgan_threshold = {\n    'Hubmap': {\n        'kidney'        : 0.40,\n        'prostate'      : 0.40,\n        'largeintestine': 0.40,\n        'spleen'        : 0.40,\n        'lung'          : 0.10,\n    },\n    'HPA': {\n        'kidney'        : 0.50,\n        'prostate'      : 0.50,\n        'largeintestine': 0.50,\n        'spleen'        : 0.50,\n        'lung'          : 0.10,\n    },\n}\n\ndata_source =['Hubmap', 'HPA']\norgan = ['kidney', 'prostate', 'largeintestine', 'spleen', 'lung']\n\n\n#data_source =['Hubmap',]\n#organ = ['spleen']\n\n\n#submit_type  = 'local-cv'    \nsubmit_type  = 'local-test'    \n#submit_type  = 'kaggle'   ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-11T16:19:38.567772Z","iopub.execute_input":"2022-08-11T16:19:38.568902Z","iopub.status.idle":"2022-08-11T16:19:38.578968Z","shell.execute_reply.started":"2022-08-11T16:19:38.568853Z","shell.execute_reply":"2022-08-11T16:19:38.577910Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n## model ##################################\nfrom coat import *\nfrom daformer import *\n\nmodel = [\n    dotdict(\n        is_use = 1,\n        module = 'model_daformer_coat',\n        param={'encoder': coat_lite_medium, 'decoder':daformer_conv3x3},\n        checkpoint = [\n            '../input/hubmap-submit-06-weight0/daformer_conv3x3-coat_lite_medium-aug5b-768-fold-3-swa.pth'\n        ],\n    ),\n   \n\n]","metadata":{"execution":{"iopub.status.busy":"2022-08-11T16:19:38.711515Z","iopub.execute_input":"2022-08-11T16:19:38.712442Z","iopub.status.idle":"2022-08-11T16:19:38.718630Z","shell.execute_reply.started":"2022-08-11T16:19:38.712395Z","shell.execute_reply":"2022-08-11T16:19:38.717321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## dataset #####\n\nif submit_type == 'local-cv':\n    valid_file = '../input/hubmap-submit-06/valid_df.fold3.csv'\n    tiff_dir   = '../input/hubmap-organ-segmentation/train_images'\n\nif (submit_type == 'local-test') or (submit_type == 'kaggle'):\n    valid_file = '../input/hubmap-organ-segmentation/test.csv'\n    tiff_dir   = '../input/hubmap-organ-segmentation/test_images'\n\nvalid_df = pd.read_csv(valid_file)\nvalid_df.loc[:,'img_area']=valid_df['img_height']*valid_df['img_width']#sort by biggest image first for memory debug\nvalid_df = valid_df.sort_values('img_area').reset_index(drop=True)\nprint('load valid_df ok')\n\n\ndef image_to_tensor(image, mode='rgb'):\n    if  mode=='bgr' :\n        image = image[:,:,::-1]\n    \n    x = image.transpose(2,0,1)\n    x = np.ascontiguousarray(x)\n    x = torch.tensor(x)\n    return x","metadata":{"execution":{"iopub.status.busy":"2022-08-11T16:19:38.790204Z","iopub.execute_input":"2022-08-11T16:19:38.790887Z","iopub.status.idle":"2022-08-11T16:19:38.814401Z","shell.execute_reply.started":"2022-08-11T16:19:38.790851Z","shell.execute_reply":"2022-08-11T16:19:38.813253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## submission ##\n\ndef do_local_validation():\n    print('\\tlocal validation ...')\n    \n    submit_df = pd.read_csv('submission.csv').fillna('')\n    submit_df = submit_df.sort_values('id')\n    truth_df  = valid_df.sort_values('id')\n    \n    lb_score = []\n    num = len(submit_df)\n    for i in range(num):\n        t_df = truth_df.iloc[i]\n        p_df = submit_df.iloc[i]\n        t = rle_decode(t_df.rle, t_df.img_height, t_df.img_width, 1)\n        p = rle_decode(p_df.rle, t_df.img_height, t_df.img_width, 1)\n        \n        dice = 2*(t*p).sum()/(p.sum()+t.sum())\n        lb_score.append(dice)\n        \n        if 0:\n            overlay = result_to_overlay(p, t)\n            image_show_norm('overlay', overlay, min=0, max=1, resize=0.10)\n            cv2.waitKey(1)\n\n    truth_df.loc[:,'lb_score']=lb_score\n    for organ in ['all', 'kidney', 'prostate', 'largeintestine', 'spleen', 'lung']:\n        if organ != 'all':\n            d = truth_df[truth_df.organ == organ]\n        else:\n            d = truth_df\n        print('\\t%f\\t%s\\t%f' % (len(d) / len(truth_df), organ, d.lb_score.mean()))\n        \n    \ndef load_net(model):\n    print('\\tload %s ... '%(model.module),end='',flush=True)\n    M = importlib.import_module(model.module)\n    num = len(model.checkpoint)\n    net = []\n    for f in range(num):\n        n = M.Net(**model.param)\n        n.load_state_dict(\n            torch.load(model.checkpoint[f], map_location=lambda storage, loc: storage) ['state_dict'],\n            strict=False)\n        n.cuda()\n        n.eval()\n        net.append(n)\n        \n    print('ok!')\n    return net\n\n\ndef do_tta_batch(image, organ):\n    \n    batch = { #<todo> multiscale????\n        'image': torch.stack([\n            image,\n            torch.flip(image,dims=[1]),\n            torch.flip(image,dims=[2]),\n        ]),\n        'organ': torch.Tensor(\n            [[organ_to_label[organ]]]*3\n        ).long()\n    }\n    return batch\n\ndef undo_tta_batch(probability):\n    probability[0] = probability[0]\n    probability[1] = torch.flip(probability[1],dims=[1])\n    probability[2] = torch.flip(probability[2],dims=[2])\n    probability = probability.mean(0, keepdims=True)\n    probability = probability[0,0].float()\n    return probability\n\ndef do_submit(): \n    print('** submit_type  = %s *******************'%submit_type)\n\n    all_net = [ load_net(m) for m in model if m.is_use==1 ]\n    \n    result = []\n    start_timer = timer()\n    for i,d in valid_df.iterrows():\n        id = d['id']\n        if (d['data_source'] in data_source) and (d['organ'] in organ):\n            \n            tiff_file = tiff_dir +'/%d.tiff'%id\n            tiff = read_tiff(tiff_file, 'rgb') \n            tiff = tiff.astype(np.float32)/255\n            H,W,_ = tiff.shape\n            \n            if 0:\n                s = d.pixel_size/0.4 * (image_size/3000)\n                h = int(np.ceil(int(H*s)/32)*32)\n                w = int(np.ceil(int(W*s)/32)*32) \n                image = cv2.resize(tiff,dsize=(w,h),interpolation=cv2.INTER_LINEAR)\n            else: \n                #or just resize to h,w = 768\n                image = cv2.resize(tiff,dsize=(image_size,image_size),interpolation=cv2.INTER_LINEAR)\n            \n            image = image_to_tensor(image, 'rgb')\n            batch = { k:v.cuda() for k,v in do_tta_batch(image, d.organ).items() }\n    \n            use = 0\n            probability = 0\n            with torch.no_grad():\n                with amp.autocast(enabled = True):\n                    \n                    for net in all_net:\n                        for n in net:\n                            use += 1\n                            output = n(batch)#data_parallel(net, batch) #\n                            probability += \\\n                                F.interpolate(output['probability'], size=(d.img_height,d.img_width),\n                                              mode='bilinear',align_corners=False, antialias=True )\n                       \n                    probability = undo_tta_batch(probability/use)\n            #---\n            probability = probability.data.cpu().numpy()\n            p = probability>organ_threshold[d.data_source][d.organ] \n            rle = rle_encode(p)\n        else:\n            rle = ''\n        \n        #----\n        if 0: #debug\n            image = cv2.cvtColor(tiff, 4).astype(np.float32)/255 #cv2.COLOR_RGB2BGR=4\n            mask  = rle_decode(d.rle, d.img_height, d.img_width, 1) #None\n            overlay = result_to_overlay(image, mask, probability)\n            \n            #image_show('image',image, resize=0.25)\n            image_show('overlay',overlay, resize=0.25)\n            cv2.waitKey(0)\n            pass\n        \n        result.append({ 'id':id, 'rle':rle, })\n        print('\\r', '\\tsubmit ... %3d/%3d %s'%(i, len(valid_df), time_to_str(timer() - start_timer,'sec')), end='',flush=True)\n    print('\\n')\n    \n    #---\n    submit_df = pd.DataFrame(result)\n    submit_df.to_csv('submission.csv',index=False)\n    print(submit_df)\n    print('\\tsubmit_df ok!')\n    print('')\n    \n    if submit_type  == 'local-cv':\n        do_local_validation()\n        \n    if submit_type == 'local-test':\n        import matplotlib.pyplot as plt \n        m = tiff\n        p = probability\n        \n        plt.figure(figsize=(12, 7))\n        plt.subplot(1, 3, 1); plt.imshow(m); plt.axis('OFF'); plt.title('image')\n        plt.subplot(1, 3, 2); plt.imshow(p*255); plt.axis('OFF'); plt.title('mask')\n        plt.subplot(1, 3, 3); plt.imshow(m); plt.imshow(p*255, alpha=0.4); plt.axis('OFF'); plt.title('overlay')\n        plt.tight_layout()\n        plt.show()\ndo_submit()","metadata":{"execution":{"iopub.status.busy":"2022-08-11T16:19:38.875455Z","iopub.execute_input":"2022-08-11T16:19:38.876122Z","iopub.status.idle":"2022-08-11T16:19:43.138106Z","shell.execute_reply.started":"2022-08-11T16:19:38.876078Z","shell.execute_reply":"2022-08-11T16:19:43.137284Z"},"trusted":true},"execution_count":null,"outputs":[]}]}