{"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 work in progress   \n\n\nplease refer to https://www.kaggle.com/competitions/rsna-breast-cancer-detection/discussion/375520\n\ntodo:\n\n- update prototype in training\n- ~prototype loss (cluster and sepration loss)~\n- example of train loop (how to create prototype)\n\n[paper]\n\n[1] Knowledge Distillation to Ensemble Global and Interpretable Prototype-based Mammogram Classification Models  \nhttps://arxiv.org/abs/2209.12420\n\n[2] This Looks Like That: Deep Learning for Interpretable Image Recognition  \nhttps://arxiv.org/pdf/1806.10574.pdf\n\n\n<a href=\"https://ibb.co/cLmvYp4\"><img src=\"https://i.ibb.co/VmKNwhf/Selection-443.png\" alt=\"Selection-443\" border=\"0\"></a>\n\n\n<a href=\"https://ibb.co/GxZhfQF\"><img src=\"https://i.ibb.co/p0Hm6jn/Selection-446.png\" alt=\"Selection-446\" border=\"0\"></a>\n<a href=\"https://ibb.co/hR0wfSf\"><img src=\"https://i.ibb.co/DDdTpZp/Selection-444.png\" alt=\"Selection-444\" border=\"0\"></a>","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/timm-pytorch-image-models/pytorch-image-models-master')\nsys.path.append('/kaggle/input/einops/einops-master')\n\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom einops import rearrange\n\nfrom timm.models.efficientnet import *\n\nimport numpy as np\n\n'''\nhttps://arxiv.org/abs/2209.12420\nKnowledge Distillation to Ensemble Global and Interpretable Prototype-based Mammogram Classification Models\n\nhttps://arxiv.org/pdf/1806.10574.pdf\nThis Looks Like That: Deep Learning for Interpretable Image Recognition\n\n'''\n\n#parameters used in paper\nimage_height  = 1536\nimage_width   = 768\nnum_prototype = 400  #half is benign, half is malignant prototype\n\nfeature_dim = 128\n\n#prototype net\np_temperature = 128\np_lamda1 = 0.1\np_lamda2 = 0.1\np_margin = 10\n\n#loss\nkd_margin = 0.2\n\n\n\n\nclass Backbone(nn.Module):\n\tdef __init__(self, dim = feature_dim):\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\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\n\nclass GlobalNet(nn.Module):\n\tdef __init__(self, dim,  feature_dim):\n\t\tsuper(GlobalNet, self).__init__()\n\n\t\t# <todo> modify this from your experiment\n\t\t# try resnet block? effect net block? here we just use simple conv block ...\n\t\tself.conv = nn.Sequential(\n\t\t\tnn.Conv2d(dim, feature_dim, kernel_size=1, padding=0, bias=False),\n\t\t\tnn.BatchNorm2d(feature_dim),\n\t\t\tnn.ReLU(inplace=True),\n\t\t)\n\t\tself.cancer = nn.Linear(feature_dim, 1)\n\n\tdef forward(self, x):\n\t\tB, C, H, W = x.shape\n\n\t\tx = self.conv(x)\n\t\tx = F.adaptive_avg_pool2d(x, 1)\n\t\tx = torch.flatten(x, 1)\n\t\tcancer = self.cancer(x).reshape(-1)\n\n\t\treturn cancer\n\n#https://discuss.pytorch.org/t/how-to-do-a-unravel-index-in-pytorch-just-like-in-numpy/12987/2\ndef torch_unravel_index(index, shape):\n\tout = []\n\tfor dim in reversed(shape):\n\t\tout.append(index % dim)\n\t\tindex =  torch.div(index, dim, rounding_mode='trunc') #index // dim\n\treturn tuple(reversed(out))\n\n\n# https://github.com/cfchen-duke/ProtoPNet\n# this is the simplest for of prototype net.\n# refer to github or new paper for batter enchancements ....\nclass PrototypeNet(nn.Module):\n\n\tdef update_prototype(self):\n\t\tpass\n\n\t\t# for learning prototype\n\t\t# eqn(6)\n\n\n\tdef __init__(self, dim,  feature_dim, num_prototype=num_prototype):\n\t\tsuper(PrototypeNet, self).__init__()\n\t\tself.num_prototype = num_prototype\n\t\tself.prototype = nn.Parameter(torch.rand((num_prototype, feature_dim )), requires_grad=True)\n\n\t\tself.conv = nn.Sequential(\n\t\t\tnn.Conv2d(dim, feature_dim, kernel_size=1, padding=0, bias=False),\n\t\t\tnn.BatchNorm2d(feature_dim),\n\t\t\tnn.ReLU(inplace=True),\n\t\t)\n\n\t\tself.logit = nn.Linear(num_prototype, 1, bias=False)  # do not use bias\n\n\n\tdef forward(self, x):\n\t\tx = self.conv(x)\n\t\tbatch_size,dim,h,w = x.shape\n\n\t\t# compute l2 distance (x-p)**2\n\t\tx = x.unsqueeze(2) # [B, 1, 128, 48, 24]\n\t\tdistance = x - self.prototype.t().reshape(1,dim,self.num_prototype,1,1)\n\t\tdistance = distance ** 2\n\t\tdistance = distance.sum(1) #[B, 400, 48, 24]\n\n\t\t#<todo> check if we should use dist or torchsq(dist)\n\n\t\tmin_dist, min_index = distance.reshape( batch_size, self.num_prototype,-1).min(-1)\n\t\tmin_index = torch.stack(torch_unravel_index(min_index, (h,w)),-1)\n\n\t\t#min distance\n\t\tsimilarity = torch.exp(-min_dist/p_temperature)\n\t\tlogit = self.logit(similarity).reshape(-1)\n\n\t\treturn logit, distance, min_dist, min_index\n\n\n\n###########################################################################3\n\n\nclass Net(nn.Module):\n\tdef load_pretrain(self, ):\n\t\treturn\n\n\tdef __init__(self, dim = feature_dim ):\n\t\tsuper(Net, self).__init__()\n\t\tself.output_type = ['inference', 'loss']\n\n\t\tself.backbone = Backbone()\n\t\tbackbone_dim = 1280\n\n\t\tself.prototype_net  = PrototypeNet(backbone_dim, dim)  # local\n\t\tself.global_net  = GlobalNet(backbone_dim, dim)    # global\n\n\n\tdef forward(self, batch):\n\t\tx = batch['image']\n\t\tbatch_size,C,H,W = x.shape\n\n\t\tf = self.backbone(x)\n\t\tlocal_logit,  distance, min_dist, min_index  = self.prototype_net(f)\n\t\tlocal_probability = torch.sigmoid(local_logit)\n\n\t\tglobal_logit = self.global_net(f)\n\t\tglobal_probability = torch.sigmoid(global_logit)\n\n\t\toutput = {}\n\t\tif  'loss' in self.output_type:\n\t\t\toutput['global_ce_loss'] = F.binary_cross_entropy_with_logits(global_logit, batch['cancer'])\n\t\t\toutput['prototype_loss'] = criterion_prototype_loss(local_logit, min_dist, batch['cancer'])\n\t\t\toutput['kd_loss'] = criterion_kd_distillation(global_probability, local_probability, batch['cancer'])\n\n\n\t\tif 'inference' in self.output_type:\n\t\t\toutput['found_prototype'] =  distance, min_dist, min_index\n\t\t\toutput['global_probability'] = global_probability\n\t\t\toutput['local_probability' ] = local_probability\n\t\t\toutput['cancer'] = (local_probability + global_probability)/2\n\n\t\treturn output\n\ndef criterion_kd_distillation(global_probability, local_probability, truth):\n\n\ty_g = truth*global_probability + (1-truth)*(1-global_probability)\n\ty_l = truth*local_probability  + (1-truth)*(1-local_probability)\n\n\tloss = torch.clamp(y_g-y_l+kd_margin, min=0)\n\tloss = loss.mean()\n\treturn loss\n\n# https://github.com/cfchen-duke/ProtoPNet/blob/master/train_and_test.py\ndef criterion_prototype_loss(logit, min_dist, truth):\n\n\tcross_entropy = F.binary_cross_entropy_with_logits(logit, truth)\n\n\t# we know for num_prototype, first half is benign(class=0), next half is malignant(class=1))\n\tbatch_size, num_prototype = min_dist.shape\n\n\ty_hat = truth.long()\n\tprototype_onehot = torch.BoolTensor([[True,False]]*(num_prototype//2)+[[False,True]]*(num_prototype//2) ).to(min_dist.device)\n\n\t# cluster cost (eqn (4) of paper)\n\tcorrect_class = torch.t(prototype_onehot[:, y_hat])\n\tcorrect_class_distance = torch.stack([min_dist[b][correct_class[b]] for b in range(batch_size)])\n\tcluster_cost = torch.mean(correct_class_distance.min(-1)[0])\n\n\t# separation cost (eqn (5) of paper)\n\twrong_class = ~torch.t(prototype_onehot[:, y_hat])\n\twrong_class_distance = torch.stack([min_dist[b][wrong_class[b]] for b in range(batch_size)])\n\tseparation_cost = torch.mean(wrong_class_distance.min(-1)[0])\n\n\tloss = cross_entropy + p_lamda1*cluster_cost + p_lamda2*torch.clamp(p_margin-separation_cost, min=0)\n\treturn loss\n\n\n\ndef run_check_net():\n\n\th,w = image_height, image_width\n\tbatch_size = 4\n\n\t# ---\n\tbatch = {\n\t\t'image' : torch.from_numpy(np.random.uniform(0,1,(batch_size,1,h,w))).float(),#.cuda(),\n\t\t'cancer': torch.from_numpy(np.random.choice(2,(batch_size))).float(),#.cuda(),\n\t}\n \n\n\tnet = Net()#.cuda()\n\tnet.load_pretrain()\n\n\t#net.output_type=['inference']\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 k=='found_prototype':\n\t\t\tprint('%32s :' % 'found_prototype.distance', v[0].shape)\n\t\t\tprint('%32s :' % 'found_prototype.min_dist', v[1].shape )\n\t\t\tprint('%32s :' % 'found_prototype.min_index', v[2].shape)\n\t\t\tcontinue\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":"2023-01-02T12:07:10.044596Z","iopub.execute_input":"2023-01-02T12:07:10.045793Z","iopub.status.idle":"2023-01-02T12:07:29.771463Z","shell.execute_reply.started":"2023-01-02T12:07:10.045663Z","shell.execute_reply":"2023-01-02T12:07:29.769566Z"},"trusted":true},"execution_count":null,"outputs":[]}]}