{"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 for discussion of   \nhttps://www.kaggle.com/competitions/rsna-breast-cancer-detection/discussion/370333#2120459\n\nrefrence paper:  \n[1] COVID-19 Prognosis via Self-Supervised Representation Learning and Multi-Image Prediction - A. Sriram (facebook AI), arXiv 2020  \nhttps://github.com/facebookresearch/CovidPrognosis","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport cv2\nimport numpy as np\n\n\n#-------------------------------\nimage_height = 1536\nimage_width = 960\n\n#one hot encoding\nnum_attribute = 2 + 2 + 10 + 1 #(e.g. 2 view + 2 site_id + 10 machine_id + 1 age)\n\n\n\n#here use use nextvit as the previously trained image encoder.\n'''\nwe modified nextvit to output features for each layer so that we choose to freeze later \nto reduce memory.\nat init():\n        self.out_idx = [sum(depths[:i + 1]) for i in range(len(depths))]\n        self.layer1 = self.features[         0: out_idx[0]]\n        self.layer2 = self.features[out_idx[0]: out_idx[1]]\n        self.layer3 = self.features[out_idx[1]: out_idx[2]]\n        self.layer4 = self.features[out_idx[2]: out_idx[3]]\n'''\n\n#fake NextViT() to make this notebook run\nNORM_EPS = 1e-5\nclass NextViT(nn.Module):\n    def __init__(self, ):\n        super(NextViT, self).__init__()\n        \n        self.stem=nn.Identity()\n        self.layer1=nn.Sequential(\n            nn.AdaptiveAvgPool2d(output_size=(384, 240)),\n            nn.Conv2d(3,96,kernel_size=1)\n        )\n        self.layer2=nn.Sequential(\n            nn.AdaptiveAvgPool2d(output_size=(192, 120)),\n            nn.Conv2d(96,256,kernel_size=1)\n        )\n        self.layer3=nn.Sequential(\n            nn.AdaptiveAvgPool2d(output_size=(96, 60)),\n            nn.Conv2d(256,512,kernel_size=1)\n        )\n        self.layer4=nn.Sequential(\n            nn.AdaptiveAvgPool2d(output_size=(48, 30)),\n            nn.Conv2d(512,1024,kernel_size=1)\n        )\n    #def forward(self, x: torch.Tensor) ->  List[torch.Tensor]:\n    def forward(self, x):\n        x  = self.stem(x)\n        x1 = self.layer1(x)\n        x2 = self.layer2(x1)\n        x3 = self.layer3(x2)\n        x4 = self.layer4(x3)\n        out = [x1,x2,x3,x4]\n        return out\n    \n \n    \nclass ImageNet(nn.Module):\n\n    def __init__(self, ):\n        super(ImageNet, self).__init__()\n        self.register_buffer('mean', torch.FloatTensor([0.5, 0.5, 0.5]).reshape(1, 3, 1, 1))\n        self.register_buffer('std', torch.FloatTensor([0.5, 0.5, 0.5]).reshape(1, 3, 1, 1))\n        self.encoder = NextViT() #nextvit_base(pretrained=True)\n        self.cancer  = nn.Linear(1024,1)\n         \n\n    def forward_feature(self, x):\n        batch_size,C,H,W = x.shape\n        x = (x - self.mean) / self.std\n        encode = self.encoder.forward(x)\n        last = encode[-1]\n        return last\n\n    def forward(self, x): \n        last = self.forward_feature(x)\n        \n        #classifier head\n        last = self.encoder.norm(last)\n        last = F.adaptive_avg_pool2d(last,1)\n        last = torch.flatten(last,1,3)\n        cancer = self.cancer(last).reshape(-1)\n        cancer = torch.sigmoid(cancer)\n        return cancer\n    \n##########################################################################\n# helper\n\ndef pad_tensor(t, length):\n    batch_size = len(t)\n    dim = t[0].shape[-1]\n    max_L = max(length)\n\n    pad_t = torch.ones((batch_size, max_L, dim)).to(t[0].device)\n    pad_mask = torch.zeros((batch_size, max_L)).to(t[0].device)\n    for b in range(batch_size):\n        pad_t[b, :length[b]] = t[b]\n        pad_mask[b, :length[b]] = 0\n    pad_mask = pad_mask > 0.5\n    return pad_t, pad_mask\n\n\ndef unpad_tensor(pad_t, length):\n    batch_size = len(pad_t)\n    t = []\n    for b in range(batch_size):\n        t.append(pad_t[b,:length[b]])\n    return t\n\ndef extract_image_feature_without_grad(image_net, image):\n    #todo : speedup with tensorrt? or compiled torch sript?\n\n    num_image = len(image)\n    image_feature = []\n\n    image_net.eval()\n    with torch.no_grad():\n        with torch.cuda.amp.autocast(enabled=True):\n            for b in range(0, num_image, 8):\n                f = image_net.forward_feature(image[b:b + 8])\n                image_feature.append(f)\n    image_feature = torch.concat(image_feature)\n    return image_feature\n    \n#################################################################################\n# multi-image model\n\n\nclass Net(nn.Module):\n    def load_pretrain(self, ):\n        return #fake function to make this notebook run\n        pretain = '/.../nextvit-b-1536-fold0-swa.lb0.59.model.pth' \n        print('load %s' % pretain)\n        state_dict = torch.load(pretain, map_location=lambda storage, loc: storage)['state_dict']  # True\n        print(self.image_net.load_state_dict(state_dict, strict=False))  \n\n    def __init__(self,):\n        super(Net, self).__init__()\n        self.output_type = ['inference', 'loss']\n\n        image_dim = 1024 \n        attribute_dim = num_attribute\n        transformer_dim = 64\n\n        self.image_net = ImageNet()\n        self.norm = nn.BatchNorm2d(image_dim, eps=NORM_EPS)\n\n        self.embed_image = nn.Linear(image_dim, transformer_dim)\n        self.embed_attribute = nn.Linear(attribute_dim, transformer_dim)\n        self.transformer = nn.TransformerEncoderLayer(\n            d_model=transformer_dim,\n            dim_feedforward=2*transformer_dim,\n            nhead=4,\n            dropout=0.25,\n            batch_first=True,\n        ) \n        self.cancer = nn.Linear(image_dim+transformer_dim,1)\n\n    def forward(self, batch):\n        image = batch['image']\n        attribute = batch['attribute']\n\n        batch_size = batch['batch_size']\n        length = batch['length']\n        num_image, C, H, W = image.shape\n\n        #----\n        # image encoder\n        image_feature = extract_image_feature_without_grad(self.image_net, image)\n        image_feature = F.relu(self.norm(image_feature))\n\n        # here we use global pooling as an example.\n        # todo region pooling in future\n        f = F.adaptive_avg_pool2d(image_feature,1)\n        f = torch.flatten(f,1,3)\n\n        #----\n        # transformer\n\n        t0 = self.embed_image(f)\n        #t1 = self.embed_attribute(attribute)\n        t = t0 #torch.cat([t0, t1],-1) or t0+t1 #todo\n        t = torch.split(t, split_size_or_sections=length, dim=0)\n\n        # https://stackoverflow.com/questions/62170439/difference-between-src-mask-and-src-key-padding-mask\n        pad_t, pad_mask = pad_tensor(t, length)\n        pad_t = self.transformer(pad_t, src_mask=None, src_key_padding_mask=pad_mask)\n\n\n        #-------\n        # pool\n        t = unpad_tensor(pad_t, length)\n        f = torch.split(f, split_size_or_sections=length, dim=0)\n\n        pool = []\n        for b in range(batch_size):\n            p = torch.cat([f[b], t[b]], -1).sum(0)\n            pool.append(p)\n        pool = torch.stack(pool)\n\n        #-------\n        # classifier\n        cancer = self.cancer(pool).reshape(-1)\n\n        output = {}\n        if  'loss' in self.output_type:\n            output['cancer_loss']=F.binary_cross_entropy_with_logits(cancer,batch['cancer'])\n\n        if 'inference' in self.output_type:\n            output['cancer']=torch.sigmoid(cancer)\n\n        return output\n\n\ndef run_check_net():\n\n    h, w = image_height, image_width\n    batch_size = 4\n    length = [2,5,1,3] #first breast have 2 images, next breast have 5 images, etc ....\n    num_image = sum(length)\n \n    # dummy data\n    batch = {\n        'batch_size' : batch_size,\n        'length' : length,\n        'attribute' : torch.from_numpy(np.random.uniform(0, 1, (num_image, num_attribute))).float(),#.cuda(),\n        'image' : torch.from_numpy(np.random.uniform(0,1,(num_image, 1, h,w))).float(),#.cuda(),\n        'cancer': torch.from_numpy(np.random.choice(2,(batch_size))).float(),#.cuda(),\n    } \n\n    net = Net()#.cuda()\n    net.load_pretrain() \n\n    with torch.no_grad():\n        with torch.cuda.amp.autocast(enabled=True):\n            output = net(batch)\n\n    print('batch')\n    for k, v in batch.items():\n        if any(c in k for c in ['length','batch_size']) : continue\n        print('%32s :' % k, v.shape)\n\n    print('output')\n    for k, v in output.items():\n        if 'loss' not in k:\n            print('%32s :' % k, v.shape)\n    print('')\n    for k, v in output.items():\n        if 'loss' in k:\n            print('%32s :' % k, v.item())\n\n    \nrun_check_net()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-29T21:17:52.840977Z","iopub.execute_input":"2023-01-29T21:17:52.841509Z","iopub.status.idle":"2023-01-29T21:17:55.216308Z","shell.execute_reply.started":"2023-01-29T21:17:52.84147Z","shell.execute_reply":"2023-01-29T21:17:55.21499Z"},"trusted":true},"execution_count":null,"outputs":[]}]}