{"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":"ADMANI dataset[1] is part of this kaggle rasna, see discussion https://www.kaggle.com/competitions/rsna-breast-cancer-detection/discussion/370333#2076911.\n\nThere are a few papers that shows good results with ADMANI dataset[1], which could be used as baseline for this compeition. \n\nHowever,there are not open source implementation, hence I implement on of them [2]\n\n![https://i.ibb.co/4S9BcxG/Selection-318.png](https://i.ibb.co/4S9BcxG/Selection-318.png)\n![https://i.ibb.co/hskqt04/Selection-319.png](https://i.ibb.co/hskqt04/Selection-319.png)\n\n\n[1] ADMANI: Annotated Digital Mammograms and Associated Non-Image Datasets  \nhttps://pubs.rsna.org/doi/pdf/10.1148/ryai.220072\n\n[2] Multi-view Local Co-occurrence and Global Consistency Learning Improve Mammogram   Classification Generalisation\nhttps://arxiv.org/pdf/2209.10478.pdf\n","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/timm-pytorch-image-models/pytorch-image-models-master')\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom timm.models.efficientnet import *\n\nimport numpy as np\n\n'''\nhttps://arxiv.org/pdf/2209.10478.pdf\n[1] Multi-view Local Co-occurrence and Global Consistency Learning Improve Mammogram\nClassification Generalisation\n'''\n\nclass Backbone(nn.Module):\n\tdef __init__(self, ):\n\t\tsuper(Backbone, self).__init__()\n\t\tself.register_buffer('mean', torch.FloatTensor([0.5, 0.5, 0.5]).reshape(1,3,1,1))\n\t\tself.register_buffer('std',  torch.FloatTensor([0.5, 0.5, 0.5]).reshape(1,3,1,1))\n\t\tself.encoder = efficientnet_b0(pretrained=True, drop_rate = 0.2, drop_path_rate = 0.1)\n\n\tdef forward(self, x):\n\t\tx = (x - self.mean) / self.std\n\t\tx = self.encoder.forward_features(x)\n\t\treturn x\n\nclass GlobalConsistency(nn.Module):\n\tdef __init__(self,dim):\n\t\tsuper(GlobalConsistency, self).__init__()\n\n\t\tself.project = nn.Linear(dim,dim) #<todo> try mlp?\n\n\tdef forward(self, u_m, u_a):\n\t\tB, C, H, W = u_m.shape\n\n\t\tg_m = F.adaptive_max_pool2d(u_m,1)\n\t\tg_a = F.adaptive_max_pool2d(u_a,1)\n\t\tg_m = torch.flatten(g_m, 1)\n\t\tg_a = torch.flatten(g_a, 1)\n\n\t\tp_a = self.project(g_m)\n\t\tp_m = self.project(g_a)\n\n\t\treturn g_m, p_m, g_a, p_a\n\n#----\n#<todo>\n#do we need norm?\n\nclass CrossAttention(nn.Module):\n\tdef __init__(self, dim, num_head=8, qkv_bias=False, attn_drop=0., proj_drop=0.):\n\t\tsuper(CrossAttention, self).__init__()\n\n\t\tassert dim % num_head == 0, 'dim should be divisible by num_heads'\n\t\tself.num_head = num_head\n\t\thead_dim = dim // num_head\n\t\tself.scale = head_dim ** (-0.5)\n\n\t\tself.qkv_a = nn.Linear(dim, dim * 3, bias=qkv_bias)\n\t\tself.qkv_m = nn.Linear(dim, dim * 3, bias=qkv_bias)\n\t\tself.proj_a = nn.Linear(dim, dim)\n\t\tself.proj_m = nn.Linear(dim, dim)\n\n\t\tself.attn_drop = nn.Dropout(attn_drop)\n\t\tself.proj_drop = nn.Dropout(proj_drop)\n\n\tdef forward(self, u_m, u_a):\n\t\tB,L,dim = u_m.shape\n\n\t\tqkv_m = self.qkv_m(u_m)\n\t\tqkv_m = qkv_m.reshape(B, L, 3, self.num_head, dim // self.num_head).permute(2, 0, 3, 1, 4)\n\t\tq_m, k_m, v_m = qkv_m.unbind(0)\n\n\t\tqkv_a = self.qkv_m(u_a)\n\t\tqkv_a = qkv_a.reshape(B, L, 3, self.num_head, dim // self.num_head).permute(2, 0, 3, 1, 4)\n\t\tq_a, k_a, v_a = qkv_a.unbind(0)\n\n\t\tattn_m = (q_m @ k_a.transpose(-2, -1)) * self.scale\n\t\tattn_m = attn_m.softmax(dim=-1)\n\t\tattn_m = self.attn_drop(attn_m)\n\n\t\tattn_a = (q_a @ k_m.transpose(-2, -1)) * self.scale\n\t\tattn_a = attn_a.softmax(dim=-1)\n\t\tattn_a = self.attn_drop(attn_a)\n\n\n\t\tx_m = (attn_m @ v_m).transpose(1, 2).reshape(B, L, dim)\n\t\tx_m = self.proj_m(x_m)\n\t\tx_m = self.proj_drop(x_m)\n\n\t\tx_a = (attn_a @ v_a).transpose(1, 2).reshape(B, L, dim)\n\t\tx_a = self.proj_a(x_a)\n\t\tx_a = self.proj_drop(x_a)\n\n\t\treturn  x_m, x_a\n\nclass LocalCoccurrence(nn.Module):\n\tdef __init__(self,dim, num_head=8, qkv_bias=False, attn_drop=0., proj_drop=0.):\n\t\tsuper(LocalCoccurrence, self).__init__()\n\n\t\tself.norm1 = nn.LayerNorm(dim)\n\t\tself.attn  = CrossAttention(dim, num_head, qkv_bias, attn_drop, proj_drop)\n\n\tdef forward(self, u_m, u_a):\n\t\tB,C,H,W = u_m.shape\n\t\tL = H*W\n\t\tdim = C\n\n\t\tu_m = u_m.reshape(B,dim,L).permute(0,2,1)\n\t\tu_a = u_a.reshape(B,dim,L).permute(0,2,1)\n\n\t\tx_m = self.norm1(u_m)\n\t\tx_a = self.norm1(u_a)\n\t\tx_m, x_a = self.attn(x_m, x_a)\n\n\t\tgap_m = x_m.mean(1)\n\t\tgap_a = x_a.mean(1)\n\t\treturn gap_m, gap_a\n\n\n\n\nclass Net(nn.Module):\n\tdef load_pretrain(self, ):\n\t\treturn\n\n\tdef __init__(self,):\n\t\tsuper(Net, self).__init__()\n\t\tself.output_type = ['inference', 'loss']\n\n\n\t\tself.backbone = Backbone()\n\t\tdim = 1280\n\n\t\tself.lc  = LocalCoccurrence(dim)\n\t\tself.gl  = GlobalConsistency(dim)\n\t\tself.mlp = nn.Sequential(\n\t\t\tnn.LayerNorm(dim*3),\n\t\t\tnn.Linear(dim*3, dim),\n\t\t\tnn.GELU(),\n\t\t\tnn.Linear(dim, dim),\n\t\t)#<todo> mlp needs to be deep if backbone is strong?\n\t\tself.cancer = nn.Linear(dim,1)\n\n\tdef forward(self, batch):\n\t\tx = batch['image']\n\t\tbatch_size,num_view,C,H,W = x.shape\n\t\tx = x.reshape(-1, C, H, W)\n\n\t\tu = self.backbone(x)\n\t\t_,c,h,w = u.shape\n\n\t\tu = u.reshape(batch_size,num_view,c,h,w)\n\t\tu_m = u[:,0]\n\t\tu_a = u[:,1]\n\t\tgap_m, gap_a = self.lc(u_m, u_a)\n\n\t\tg_m, p_m, g_a, p_a = self.gl(u_m, u_a)\n\t\tgp_m = g_m + p_m\n\n\t\tlast = torch.cat([gp_m, gap_m, gap_a ],-1)\n\t\tlast = self.mlp(last)\n\t\tcancer = self.cancer(last).reshape(-1)\n\n\n\t\toutput = {}\n\t\tif  'loss' in self.output_type:\n\t\t\toutput['cancer_loss'] = F.binary_cross_entropy_with_logits(cancer, batch['cancer'])\n\t\t\toutput['global_loss'] = criterion_global_consistency(g_m, p_m, g_a, p_a)\n\n\n\t\tif 'inference' in self.output_type:\n\t\t\toutput['cancer'] = torch.sigmoid(cancer)\n\n\t\treturn output\n\ndef similiarity(x1,x2):\n\tp12 = (x1*x2).sum(-1)\n\tp1 = torch.sqrt((x1*x1).sum(-1))\n\tp2 = torch.sqrt((x2*x2).sum(-1))\n\ts = p12/(p1*p2+1e-6)\n\treturn s\n\n\ndef criterion_global_consistency(g_m, p_m, g_a, p_a):\n\tloss =  -0.5*(similiarity(g_m, p_m)+similiarity(g_a, p_a))\n\tloss = loss.mean()\n\treturn loss\n\n\n\ndef run_check_net():\n\n\th,w = 1536, 768 #about 0.50\n\tbatch_size = 4\n\n\t# ---\n\tbatch = {\n\t\t'image' : torch.from_numpy(np.random.uniform(0,1,(batch_size,2,1,h,w))).float(),#.cuda(),\n\t\t'cancer': torch.from_numpy(np.random.choice(2,(batch_size))).float(),#.cuda(),\n\t}\n\t#batch = {k: v.cuda() for k, v in batch.items()}\n\n\tnet = Net()#.cuda()\n\tnet.load_pretrain()\n\n\twith torch.no_grad():\n\t\twith torch.cuda.amp.autocast(enabled=True):\n\t\t\toutput = net(batch)\n\n\tprint('batch')\n\tfor k, v in batch.items():\n\t\tif 'index' in k: continue\n\t\tprint('%32s :' % k, v.shape)\n\n\tprint('output')\n\tfor k, v in output.items():\n\t\tif 'loss' not in k:\n\t\t\tprint('%32s :' % k, v.shape)\n\tfor k, v in output.items():\n\t\tif 'loss' in k:\n\t\t\tprint('%32s :' % k, v.item())\n\n# main #################################################################\nif __name__ == '__main__':\n\trun_check_net()","metadata":{"execution":{"iopub.status.busy":"2022-12-27T05:48:00.830956Z","iopub.execute_input":"2022-12-27T05:48:00.83138Z","iopub.status.idle":"2022-12-27T05:48:13.150746Z","shell.execute_reply.started":"2022-12-27T05:48:00.831338Z","shell.execute_reply":"2022-12-27T05:48:13.14942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"there is starter code of baseline pure efficient-b0 single view here:\nhttps://www.kaggle.com/datasets/hengck23/pure-effb0-single-view-starter-for-mvccl-admani\n\n\n<a href=\"https://ibb.co/bWMW256\"><img src=\"https://i.ibb.co/X2r23xW/Selection-469.png\" alt=\"Selection-469\" border=\"0\"></a>\n<a href=\"https://ibb.co/8rWj273\"><img src=\"https://i.ibb.co/n7HDfL2/Selection-470.png\" alt=\"Selection-470\" border=\"0\"></a>\n","metadata":{}}]}