{"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"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":20270,"databundleVersionId":1222630,"sourceType":"competition"},{"sourceId":1243687,"sourceType":"datasetVersion","datasetId":690737},{"sourceId":1253590,"sourceType":"datasetVersion","datasetId":720563}],"dockerImageVersionId":29928,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nfrom tqdm import tqdm\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader,Dataset\nfrom torchvision import datasets, models,transforms \nimport torchvision.utils as utils\nimport os\nfrom sklearn.model_selection import train_test_split\nfrom PIL import Image\nimport os\nfrom sklearn.metrics import auc,roc_auc_score\ndevice = torch.device(\"cpu\")\nimport warnings\nwarnings.filterwarnings(\"ignore\")\nimport matplotlib.pyplot as plt\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport random\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-06-09T10:26:51.376849Z","iopub.execute_input":"2022-06-09T10:26:51.377188Z","iopub.status.idle":"2022-06-09T10:26:51.385739Z","shell.execute_reply.started":"2022-06-09T10:26:51.377160Z","shell.execute_reply":"2022-06-09T10:26:51.384484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    \"\"\"\n    Seeds basic parameters for reproductibility of results\n    \"\"\"\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    torch.backends.cudnn.benchmark = False\n\nseed_everything(42)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-06-09T10:26:52.060170Z","iopub.execute_input":"2022-06-09T10:26:52.060500Z","iopub.status.idle":"2022-06-09T10:26:52.069361Z","shell.execute_reply.started":"2022-06-09T10:26:52.060468Z","shell.execute_reply":"2022-06-09T10:26:52.068621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# read the data \ndata_dir='../input/melanoma-merged-external-data-512x512-jpeg/512x512-dataset-melanoma/512x512-dataset-melanoma/'\ndata=pd.read_csv('../input/melanoma-merged-external-data-512x512-jpeg/marking.csv')\n\n# balance the data a bit\ndf_0=data[data['target']==0].sample(6000,random_state=42)\ndf_1=data[data['target']==1]\ndata=pd.concat([df_0,df_1]).reset_index()\n\n#prepare the data\nlabels=[]\nimages=[]\nfor i in range(data.shape[0]):\n    images.append(data_dir + data['image_id'].iloc[i]+'.jpg')\n    labels.append(data['target'].iloc[i])\ndf=pd.DataFrame(images)\ndf.columns=['images']\ndf['target']=labels\n\n# Split train into train and val\nX_train, X_val, y_train, y_val = train_test_split(df['images'],df['target'], test_size=0.2, random_state=1234)","metadata":{"execution":{"iopub.status.busy":"2022-06-09T09:57:26.432286Z","iopub.execute_input":"2022-06-09T09:57:26.432600Z","iopub.status.idle":"2022-06-09T09:57:26.735566Z","shell.execute_reply.started":"2022-06-09T09:57:26.432568Z","shell.execute_reply":"2022-06-09T09:57:26.734867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = transforms.Compose([\n        transforms.RandomRotation(10),      # rotate +/- 10 degrees\n        transforms.RandomHorizontalFlip(),  # reverse 50% of images\n        transforms.Resize(224),             # resize shortest side to 224 pixels\n        transforms.CenterCrop(224),         # crop longest side to 224 pixels at center\n        transforms.ToTensor(),\n        transforms.Normalize([0.485, 0.456, 0.406],\n                             [0.229, 0.224, 0.225])\n    ])\n\nval_transform = transforms.Compose([\n        transforms.Resize(224),\n        transforms.CenterCrop(224),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485, 0.456, 0.406],\n                             [0.229, 0.224, 0.225])\n    ])","metadata":{"execution":{"iopub.status.busy":"2022-06-09T09:59:14.429537Z","iopub.execute_input":"2022-06-09T09:59:14.429889Z","iopub.status.idle":"2022-06-09T09:59:14.436857Z","shell.execute_reply.started":"2022-06-09T09:59:14.429856Z","shell.execute_reply":"2022-06-09T09:59:14.435927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImageDataset(Dataset):\n    \n    def __init__(self,data_paths,labels,transform=None,mode='train'):\n         self.data=data_paths\n         self.labels=labels\n         self.transform=transform\n         self.mode=mode\n    def __len__(self):\n       return len(self.data)\n    \n    def __getitem__(self,idx):\n        img_name = self.data[idx]\n        img = cv2.imread(img_name)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img=Image.fromarray(img)\n        if self.transform is not None:\n          img = self.transform(img)\n        img=img.cuda()\n        \n        if self.mode=='test':\n            return img,img_name\n        else:\n           \n            labels = torch.tensor(self.labels[idx]).cuda()\n\n            return img, labels","metadata":{"execution":{"iopub.status.busy":"2022-06-09T09:58:06.063049Z","iopub.execute_input":"2022-06-09T09:58:06.063318Z","iopub.status.idle":"2022-06-09T09:58:06.072179Z","shell.execute_reply.started":"2022-06-09T09:58:06.063292Z","shell.execute_reply":"2022-06-09T09:58:06.071384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset=ImageDataset(data_paths=X_train.values,labels=y_train.values,transform=train_transform)\nval_dataset=ImageDataset(data_paths=X_val.values,labels=y_val.values,transform=val_transform)\n\ntrain_loader=DataLoader(train_dataset,batch_size=100,shuffle=True)\nval_loader=DataLoader(val_dataset,batch_size=50,shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-06-09T09:59:17.708956Z","iopub.execute_input":"2022-06-09T09:59:17.709281Z","iopub.status.idle":"2022-06-09T09:59:17.714443Z","shell.execute_reply.started":"2022-06-09T09:59:17.709250Z","shell.execute_reply":"2022-06-09T09:59:17.713554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AttentionBlock(nn.Module):\n    def __init__(self, in_features_l, in_features_g, attn_features, up_factor, normalize_attn=True):\n        super(AttentionBlock, self).__init__()\n        self.up_factor = up_factor\n        self.normalize_attn = normalize_attn\n        self.W_l = nn.Conv2d(in_channels=in_features_l, out_channels=attn_features, kernel_size=1, padding=0, bias=False)\n        self.W_g = nn.Conv2d(in_channels=in_features_g, out_channels=attn_features, kernel_size=1, padding=0, bias=False)\n        self.phi = nn.Conv2d(in_channels=attn_features, out_channels=1, kernel_size=1, padding=0, bias=True)\n    def forward(self, l, g):\n        N, C, W, H = l.size()\n        l_ = self.W_l(l)\n        g_ = self.W_g(g)\n        if self.up_factor > 1:\n            g_ = F.interpolate(g_, scale_factor=self.up_factor, mode='bilinear', align_corners=False)\n        c = self.phi(F.relu(l_ + g_)) # batch_sizex1xWxH\n        \n        # compute attn map\n        if self.normalize_attn:\n            a = F.softmax(c.view(N,1,-1), dim=2).view(N,1,W,H)\n        else:\n            a = torch.sigmoid(c)\n        # re-weight the local feature\n        f = torch.mul(a.expand_as(l), l) # batch_sizexCxWxH\n        if self.normalize_attn:\n            output = f.view(N,C,-1).sum(dim=2) # weighted sum\n        else:\n            output = F.adaptive_avg_pool2d(f, (1,1)).view(N,C) # global average pooling\n        return a, output\n","metadata":{"execution":{"iopub.status.busy":"2022-06-09T10:27:17.415576Z","iopub.execute_input":"2022-06-09T10:27:17.415950Z","iopub.status.idle":"2022-06-09T10:27:17.431888Z","shell.execute_reply.started":"2022-06-09T10:27:17.415918Z","shell.execute_reply":"2022-06-09T10:27:17.431030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AttnVGG(nn.Module):\n    def __init__(self, num_classes, normalize_attn=False, dropout=None):\n        super(AttnVGG, self).__init__()\n        net = models.vgg16_bn(pretrained=True)\n        self.conv_block1 = nn.Sequential(*list(net.features.children())[0:6])\n        self.conv_block2 = nn.Sequential(*list(net.features.children())[7:13])\n        self.conv_block3 = nn.Sequential(*list(net.features.children())[14:23])\n        self.conv_block4 = nn.Sequential(*list(net.features.children())[24:33])\n        self.conv_block5 = nn.Sequential(*list(net.features.children())[34:43])\n        self.pool = nn.AvgPool2d(7, stride=1)\n        self.dpt = None\n        if dropout is not None:\n            self.dpt = nn.Dropout(dropout)\n        self.cls = nn.Linear(in_features=512+512+256, out_features=num_classes, bias=True)\n        \n       # initialize the attention blocks defined above\n        self.attn1 = AttentionBlock(256, 512, 256, 4, normalize_attn=normalize_attn)\n        self.attn2 = AttentionBlock(512, 512, 256, 2, normalize_attn=normalize_attn)\n        \n       \n        self.reset_parameters(self.cls)\n        self.reset_parameters(self.attn1)\n        self.reset_parameters(self.attn2)\n    def reset_parameters(self, module):\n        for m in module.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_in', nonlinearity='relu')\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0.)\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.constant_(m.weight, 1.)\n                nn.init.constant_(m.bias, 0.)\n            elif isinstance(m, nn.Linear):\n                nn.init.normal_(m.weight, 0., 0.01)\n                nn.init.constant_(m.bias, 0.)\n    def forward(self, x):\n        block1 = self.conv_block1(x)       # /1\n        pool1 = F.max_pool2d(block1, 2, 2) # /2\n        block2 = self.conv_block2(pool1)   # /2\n        pool2 = F.max_pool2d(block2, 2, 2) # /4\n        block3 = self.conv_block3(pool2)   # /4\n        pool3 = F.max_pool2d(block3, 2, 2) # /8\n        block4 = self.conv_block4(pool3)   # /8\n        pool4 = F.max_pool2d(block4, 2, 2) # /16\n        block5 = self.conv_block5(pool4)   # /16\n        pool5 = F.max_pool2d(block5, 2, 2) # /32\n        N, __, __, __ = pool5.size()\n        \n        g = self.pool(pool5).view(N,512)\n        a1, g1 = self.attn1(pool3, pool5)\n        a2, g2 = self.attn2(pool4, pool5)\n        g_hat = torch.cat((g,g1,g2), dim=1) # batch_size x C\n        if self.dpt is not None:\n            g_hat = self.dpt(g_hat)\n        out = self.cls(g_hat)\n\n        return [out, a1, a2]","metadata":{"execution":{"iopub.status.busy":"2022-06-09T10:27:22.016776Z","iopub.execute_input":"2022-06-09T10:27:22.017091Z","iopub.status.idle":"2022-06-09T10:27:22.042616Z","shell.execute_reply.started":"2022-06-09T10:27:22.017061Z","shell.execute_reply":"2022-06-09T10:27:22.041790Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = AttnVGG(num_classes=1, normalize_attn=True)\nmodel=model.cuda()","metadata":{"execution":{"iopub.status.busy":"2022-06-09T10:27:29.261278Z","iopub.execute_input":"2022-06-09T10:27:29.261599Z","iopub.status.idle":"2022-06-09T10:27:58.544496Z","shell.execute_reply.started":"2022-06-09T10:27:29.261568Z","shell.execute_reply":"2022-06-09T10:27:58.543575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, alpha=0.25, gamma=2.0, logits=False, reduce=True):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.logits = logits\n        self.reduce = reduce\n\n    def forward(self, inputs, targets):\n        if self.logits:\n            BCE_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduce=False)\n        else:\n            BCE_loss = F.binary_cross_entropy(inputs, targets, reduce=False)\n        pt = torch.exp(-BCE_loss)\n        F_loss = self.alpha * (1-pt)**self.gamma * BCE_loss\n\n        if self.reduce:\n            return torch.mean(F_loss)\n        else:\n            return F_loss\n        \ncriterion = FocalLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)","metadata":{"execution":{"iopub.status.busy":"2022-06-09T10:34:45.520266Z","iopub.execute_input":"2022-06-09T10:34:45.520591Z","iopub.status.idle":"2022-06-09T10:34:45.531076Z","shell.execute_reply.started":"2022-06-09T10:34:45.520560Z","shell.execute_reply":"2022-06-09T10:34:45.529922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\nstart_time = time.time()\n\nepochs = 2\n\ntrain_losses = []\ntrain_auc=[]\nval_auc=[]\n\nfor i in range(epochs):\n\n    train_preds=[]\n    train_targets=[]\n    auc_train=[]\n    loss_epoch_train=[]\n    # Run the training batches\n    for b, (X_train, y_train) in tqdm(enumerate(train_loader),total=len(train_loader)):\n        \n        b+=1\n        y_pred,_,_=model(X_train)\n        loss = criterion(torch.sigmoid(y_pred.type(torch.FloatTensor)), y_train.type(torch.FloatTensor))   \n        loss_epoch_train.append(loss.item())\n        # For plotting purpose\n        if (i==1):\n            if (b==19):\n                I_train = utils.make_grid(X_train[0:8,:,:,:], nrow=8, normalize=True, scale_each=True)\n                __, a1, a2 = model(X_train[0:8,:,:,:])\n                \n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n                 \n    try:\n        auc_train=roc_auc_score(y_train.detach().to(device).numpy(),torch.sigmoid(y_pred).detach().to(device).numpy())\n    except:\n        auc_train=0\n    train_losses.append(np.mean(loss_epoch_train))\n    train_auc.append(auc_train)\n    print(f'epoch: {i:2}   loss: {np.mean(loss_epoch_train):10.8f} AUC  : {auc_train:10.8f} ')\n    # Run the testing batches\n    \n    with torch.no_grad():\n        for b, (X_test, y_test) in enumerate(val_loader):\n            \n            y_val,_,_ = model(X_test)\n            loss = criterion(torch.sigmoid(y_val.type(torch.FloatTensor)), y_test.type(torch.FloatTensor))\n            loss_epoch_test.append(loss.item())\n    val_auc.append(auc_val)\n    print(f'Epoch: {i} Val Loss: {np.mean(loss_epoch_test):10.8f} AUC: {auc_val:10.8f} ')\n    \nprint(f'\\nDuration: {time.time() - start_time:.0f} seconds') # print the time elapsed","metadata":{"execution":{"iopub.status.busy":"2022-06-09T10:35:36.234928Z","iopub.execute_input":"2022-06-09T10:35:36.235306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualizing Attention","metadata":{}},{"cell_type":"code","source":"def visualize_attention(I_train,a,up_factor,no_attention=False):\n    img = I_train.permute((1,2,0)).cpu().numpy()\n    # compute the heatmap\n    if up_factor > 1:\n        a = F.interpolate(a, scale_factor=up_factor, mode='bilinear', align_corners=False)\n    attn = utils.make_grid(a, nrow=8, normalize=True, scale_each=True)\n    attn = attn.permute((1,2,0)).mul(255).byte().cpu().numpy()\n    attn = cv2.applyColorMap(attn, cv2.COLORMAP_JET)\n    attn = cv2.cvtColor(attn, cv2.COLOR_BGR2RGB)\n    attn = np.float32(attn) / 255\n    # add the heatmap to the image\n    img=cv2.resize(img,(466,60))\n    if no_attention:\n        return torch.from_numpy(img)\n    else:\n        vis = 0.6 * img + 0.4 * attn\n        return torch.from_numpy(vis)\n        ","metadata":{"execution":{"iopub.status.busy":"2022-06-09T11:01:27.916659Z","iopub.execute_input":"2022-06-09T11:01:27.917066Z","iopub.status.idle":"2022-06-09T11:01:27.927211Z","shell.execute_reply.started":"2022-06-09T11:01:27.917030Z","shell.execute_reply":"2022-06-09T11:01:27.926101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"orig=visualize_attention(I_train,a1,up_factor=2,no_attention=True)\nfirst=visualize_attention(I_train,a1,up_factor=2,no_attention=False)\nsecond=visualize_attention(I_train,a2,up_factor=4,no_attention=False)\n\nfig, (ax1, ax2,ax3) = plt.subplots(3, 1,figsize=(10, 10))\nax1.imshow(orig)\nax2.imshow(first)\nax3.imshow(second)\nax1.title.set_text('Input Images')\nax2.title.set_text('pool-3 attention')\nax3.title.set_text('pool-4 attention')","metadata":{"execution":{"iopub.status.busy":"2022-06-09T11:01:31.714147Z","iopub.execute_input":"2022-06-09T11:01:31.714475Z","iopub.status.idle":"2022-06-09T11:01:32.187214Z","shell.execute_reply.started":"2022-06-09T11:01:31.714444Z","shell.execute_reply":"2022-06-09T11:01:32.185803Z"},"trusted":true},"execution_count":null,"outputs":[]}]}