{"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":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-20T01:09:39.370352Z","iopub.execute_input":"2022-07-20T01:09:39.371097Z","iopub.status.idle":"2022-07-20T01:09:39.377227Z","shell.execute_reply.started":"2022-07-20T01:09:39.371048Z","shell.execute_reply":"2022-07-20T01:09:39.376266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %% Imports\nimport torch\nimport torch.nn as nn\nimport torchvision\nimport torchvision.models as models\nimport torchvision.transforms as transforms\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport torch.optim as optim\nimport copy\nimport pandas as pd \nimport sys\nsys.path.append('../input/monai-v081/')\nfrom monai.utils import first\n\nfrom monai.transforms import (\n    Compose,\n    LoadImaged,\n    AddChanneld,\n    Resized,\n    EnsureTyped,\n    Lambdad,\n    ToTensord, \n    CropForegroundd,\n    RandAffined,\n    RandFlipd,\n    RandScaleIntensityd,\n    RandShiftIntensityd,\n    RandGridDistortiond\n)\n\nfrom monai.data import CacheDataset, Dataset, DataLoader, ThreadDataLoader, decollate_batch\n\nfrom matplotlib import pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:19:39.747762Z","iopub.execute_input":"2022-07-20T01:19:39.748119Z","iopub.status.idle":"2022-07-20T01:19:39.756401Z","shell.execute_reply.started":"2022-07-20T01:19:39.748091Z","shell.execute_reply":"2022-07-20T01:19:39.755163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set GPU device\nprint(torch.cuda.is_available())\ndevice = torch.device(\"cuda:0\")","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:09:39.916346Z","iopub.execute_input":"2022-07-20T01:09:39.917001Z","iopub.status.idle":"2022-07-20T01:09:39.922422Z","shell.execute_reply.started":"2022-07-20T01:09:39.916964Z","shell.execute_reply":"2022-07-20T01:09:39.921326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %% Load data\nTRAIN_ROOT = \"../input/brain-tumor-classification-mri/Training\"\nTEST_ROOT = \"../input/brain-tumor-classification-mri/Testing\"\ntrain_dataset = torchvision.datasets.ImageFolder(root=TRAIN_ROOT)\ntest_dataset = torchvision.datasets.ImageFolder(root=TRAIN_ROOT)","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:09:40.096001Z","iopub.execute_input":"2022-07-20T01:09:40.096650Z","iopub.status.idle":"2022-07-20T01:09:40.136085Z","shell.execute_reply.started":"2022-07-20T01:09:40.096614Z","shell.execute_reply":"2022-07-20T01:09:40.135198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_metadata(row):\n    data = row['id'].split('_')\n    case = int(data[0].replace('case',''))\n    day = int(data[1].replace('day',''))\n    slice_ = int(data[-1])\n    row['case'] = case\n    row['day'] = day\n    row['slice'] = slice_\n    return row","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:09:40.294097Z","iopub.execute_input":"2022-07-20T01:09:40.294447Z","iopub.status.idle":"2022-07-20T01:09:40.300578Z","shell.execute_reply.started":"2022-07-20T01:09:40.294415Z","shell.execute_reply":"2022-07-20T01:09:40.299382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.notebook import tqdm\ntqdm.pandas()\n\ndf = pd.read_csv('../input/uw-madison-gi-tract-image-segmentation/train.csv')\ndf = df.progress_apply(get_metadata, axis=1)\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:09:40.478071Z","iopub.execute_input":"2022-07-20T01:09:40.479080Z","iopub.status.idle":"2022-07-20T01:12:35.700507Z","shell.execute_reply.started":"2022-07-20T01:09:40.479031Z","shell.execute_reply":"2022-07-20T01:12:35.699467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def path2info(row):\n    path = row['image_path']\n    data = path.split('/')\n    slice_ = int(data[-1].split('_')[1])\n    case = int(data[-3].split('_')[0].replace('case',''))\n    day = int(data[-3].split('_')[1].replace('day',''))\n    width = int(data[-1].split('_')[2])\n    height = int(data[-1].split('_')[3])\n    row['height'] = height\n    row['width'] = width\n    row['case'] = case\n    row['day'] = day\n    row['slice'] = slice_\n    return row","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:12:35.702462Z","iopub.execute_input":"2022-07-20T01:12:35.702907Z","iopub.status.idle":"2022-07-20T01:12:35.710452Z","shell.execute_reply.started":"2022-07-20T01:12:35.702872Z","shell.execute_reply":"2022-07-20T01:12:35.709332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from glob import glob\n\npaths = glob('../input/uw-madison-gi-tract-image-segmentation/train/*/*/*/*')\npath_df = pd.DataFrame(paths, columns=['image_path'])\npath_df = path_df.progress_apply(path2info, axis=1)\ndf = df.merge(path_df, on=['case','day','slice'])\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:12:35.712347Z","iopub.execute_input":"2022-07-20T01:12:35.712733Z","iopub.status.idle":"2022-07-20T01:14:19.535007Z","shell.execute_reply.started":"2022-07-20T01:12:35.712696Z","shell.execute_reply":"2022-07-20T01:14:19.533949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\n\n\ndef isNaN(num):\n    return num != num","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:19.537650Z","iopub.execute_input":"2022-07-20T01:14:19.538094Z","iopub.status.idle":"2022-07-20T01:14:19.543632Z","shell.execute_reply.started":"2022-07-20T01:14:19.538056Z","shell.execute_reply":"2022-07-20T01:14:19.542623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset\nfrom PIL import Image\nimport numpy\n\nclass GITractData(Dataset):\n    def __init__(self, df, transform = None):\n        self.samples = []\n\n        img_list = df['image_path'].to_list()\n        img_list = list(dict.fromkeys(img_list))\n  \n        label_list= []\n        \n        for seg in df['segmentation']:\n            if isNaN(seg):\n                label_list.append(0)\n            else:\n                label_list.append(1)\n        \n        labels=[]\n        \n        for i in range(0, len(label_list), 3):\n            label_set = [label_list[i], label_list[i+1], label_list[i+2]]\n            label_tensor = torch.Tensor(label_set)\n            labels.append(label_tensor)\n    \n        self.samples = [{'image': image_name, 'label': label_name} for image_name, label_name in zip(img_list, labels)]\n        \n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        return self.samples[idx]","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:19.545340Z","iopub.execute_input":"2022-07-20T01:14:19.545697Z","iopub.status.idle":"2022-07-20T01:14:19.557915Z","shell.execute_reply.started":"2022-07-20T01:14:19.545663Z","shell.execute_reply":"2022-07-20T01:14:19.556951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = GITractData(df)\nprint(dataset[:4])","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:19.561063Z","iopub.execute_input":"2022-07-20T01:14:19.561305Z","iopub.status.idle":"2022-07-20T01:14:19.769001Z","shell.execute_reply.started":"2022-07-20T01:14:19.561283Z","shell.execute_reply":"2022-07-20T01:14:19.767199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(dataset.__getitem__(3))","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:19.770333Z","iopub.execute_input":"2022-07-20T01:14:19.770780Z","iopub.status.idle":"2022-07-20T01:14:19.778400Z","shell.execute_reply.started":"2022-07-20T01:14:19.770741Z","shell.execute_reply":"2022-07-20T01:14:19.777259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nimg_list = df['image_path'].to_list()\nimg_list = list(dict.fromkeys(img_list))\nprint(img_list[:4])\n'''","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:19.780103Z","iopub.execute_input":"2022-07-20T01:14:19.780894Z","iopub.status.idle":"2022-07-20T01:14:19.788139Z","shell.execute_reply.started":"2022-07-20T01:14:19.780858Z","shell.execute_reply":"2022-07-20T01:14:19.787034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nlabel_list = []\n\nfor seg in df['segmentation']:\n    if isNaN(seg):\n        label_list.append(0)\n    else:\n        label_list.append(1)\n        \nlabels=[]\nfor i in range(0, len(label_list), 3):\n    label_set = [label_list[i], label_list[i+1], label_list[i+2]]\n    labels.append(label_set)\n\n#print(len(labels))\n#print(len(img_list))\n\ntrain_files = [{'image': image_name, 'label': label_name} for image_name, label_name in zip(img_list, labels)]\n#print(train_files[:100])\n'''","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:19.789666Z","iopub.execute_input":"2022-07-20T01:14:19.790738Z","iopub.status.idle":"2022-07-20T01:14:19.798488Z","shell.execute_reply.started":"2022-07-20T01:14:19.790697Z","shell.execute_reply":"2022-07-20T01:14:19.797347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %% Building the model\nclass CNNModel(nn.Module):\n    def __init__(self):\n        super(CNNModel, self).__init__()\n        self.vgg16 = models.vgg16(pretrained=True) \n\n        # Replace output layer according to our problem\n        in_feats = self.vgg16.classifier[6].in_features \n        self.vgg16.classifier[6] = nn.Linear(in_feats, 3)\n\n    def forward(self, x):\n        x = self.vgg16(x)\n        return x\n\nmodel = CNNModel()\nmodel.to(device)\nmodel","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:19.804921Z","iopub.execute_input":"2022-07-20T01:14:19.805651Z","iopub.status.idle":"2022-07-20T01:14:51.427483Z","shell.execute_reply.started":"2022-07-20T01:14:19.805614Z","shell.execute_reply":"2022-07-20T01:14:51.426499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model.load_state_dict(torch.load(\"../input/weights-2-epochs/best_metric_model (15).pth\"))","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:51.429060Z","iopub.execute_input":"2022-07-20T01:14:51.429713Z","iopub.status.idle":"2022-07-20T01:14:56.715660Z","shell.execute_reply.started":"2022-07-20T01:14:51.429675Z","shell.execute_reply":"2022-07-20T01:14:56.714767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\n# %% Prepare data for pretrained model\ntrain_dataset = torchvision.datasets.ImageFolder(\n        root=TRAIN_ROOT,\n        transform=transforms.Compose([\n                      transforms.Resize((255,255)),\n                      transforms.ToTensor()\n        ])\n)\n\ntest_dataset = torchvision.datasets.ImageFolder(\n        root=TEST_ROOT,\n        transform=transforms.Compose([\n                      transforms.Resize((255,255)),\n                      transforms.ToTensor()\n        ])\n)\n#train_dataset[0][0].permute(1,2,0)\n'''","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:56.716965Z","iopub.execute_input":"2022-07-20T01:14:56.717311Z","iopub.status.idle":"2022-07-20T01:14:56.727510Z","shell.execute_reply.started":"2022-07-20T01:14:56.717276Z","shell.execute_reply":"2022-07-20T01:14:56.726535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:56.728579Z","iopub.execute_input":"2022-07-20T01:14:56.729628Z","iopub.status.idle":"2022-07-20T01:14:57.925461Z","shell.execute_reply.started":"2022-07-20T01:14:56.729577Z","shell.execute_reply":"2022-07-20T01:14:57.924498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = Compose(\n    [\n        LoadImaged(keys=['image']),\n        AddChanneld(keys=['image']),\n        ToTensord(keys='image'),\n        Lambdad(keys='image', func = lambda x: x / x.max()),\n        Lambdad(keys = 'image', func = lambda x: x.repeat(3,1,1)),\n        CropForegroundd(keys = 'image', source_key = 'image'),\n        Resized(keys=['image'], spatial_size=[250,250]),\n        RandFlipd(keys='image', prob=0.5, spatial_axis=[0]),\n        RandFlipd(keys='image', prob=0.5, spatial_axis=[1]),\n        RandAffined(\n            keys = 'image',\n            prob = 0.5,\n            rotate_range=np.pi/12,\n            translate_range=(250*0.0625, 250*0.0625),\n            scale_range=(0.1,0.1),\n            mode=\"nearest\",\n            padding_mode=\"reflection\",\n        ),\n        RandGridDistortiond(keys='image', prob=0.5, distort_limit=(-0.05, 0.05), mode='nearest', padding_mode='reflection'),\n        RandScaleIntensityd(keys='image', factors=(-0.2,0.2), prob=0.5),\n        RandShiftIntensityd(keys='image', offsets=(-0.1,0.1), prob=0.5),\n        EnsureTyped(keys=['image']),\n    ]\n)\ntrain_dataset = CacheDataset(\n    data=dataset,\n    transform=train_transforms,\n    cache_rate=0.0,\n    num_workers=2,\n    copy_cache=False)\n\ntrain_loader = ThreadDataLoader(\n    train_dataset,\n    num_workers=0,\n    batch_size=32,\n    shuffle=True\n)\nprint(len(train_loader))","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:20:25.442539Z","iopub.execute_input":"2022-07-20T01:20:25.442898Z","iopub.status.idle":"2022-07-20T01:20:25.466830Z","shell.execute_reply.started":"2022-07-20T01:20:25.442867Z","shell.execute_reply":"2022-07-20T01:20:25.465818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample = first(train_loader)","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:20:29.115480Z","iopub.execute_input":"2022-07-20T01:20:29.116472Z","iopub.status.idle":"2022-07-20T01:20:30.426453Z","shell.execute_reply.started":"2022-07-20T01:20:29.116438Z","shell.execute_reply":"2022-07-20T01:20:30.425486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\n# %% Create data loaders\nbatch_size = 32\ntrain_loader = torch.utils.data.DataLoader(\n    train_dataset,\n    batch_size=batch_size,\n    shuffle=True\n)\ntest_loader = torch.utils.data.DataLoader(\n    test_dataset,\n    batch_size=batch_size,\n    shuffle=True\n)\n\nprint(len(train_loader))\n'''","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:57.945633Z","iopub.status.idle":"2022-07-20T01:14:57.946910Z","shell.execute_reply.started":"2022-07-20T01:14:57.946612Z","shell.execute_reply":"2022-07-20T01:14:57.946638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nfor i in range(32):\n    print(sample[0].shape)\n    print(sample[1][i])\n'''    \nprint(sample['image'][0].shape)\nplt.figure(\"image\", (32,16))\nfor i in range(32):\n    plt.subplot(4,8,i+1)\n    plt.title(f\"label: {sample['label'][i]}\")\n    plt.imshow(sample['image'][i,0,:,:])","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:20:40.569973Z","iopub.execute_input":"2022-07-20T01:20:40.570617Z","iopub.status.idle":"2022-07-20T01:20:43.917792Z","shell.execute_reply.started":"2022-07-20T01:20:40.570560Z","shell.execute_reply":"2022-07-20T01:20:43.915269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\n# %% Train\ncross_entropy_loss = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=0.00001)\nepochs = 2\n\n# Iterate x epochs over the train data\nfor epoch in range(epochs):  \n    for i, batch in enumerate(train_loader, 0):\n        inputs, labels = batch\n        inputs = inputs.to(device)\n        labels = labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        # Labels are automatically one-hot-encoded\n        loss = cross_entropy_loss(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        print(loss)\n    print(\"done epoch:\", epoch)\n'''","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:57.953161Z","iopub.status.idle":"2022-07-20T01:14:57.953885Z","shell.execute_reply.started":"2022-07-20T01:14:57.953629Z","shell.execute_reply":"2022-07-20T01:14:57.953652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %% Train\nbce_loss = nn.BCELoss()\noptimizer = optim.Adam(model.parameters(), lr=0.00001)\nepochs = 10\nmodel_dir = './'\n\n\n# Iterate x epochs over the train data\nfor epoch in range(epochs):  \n    for i, batch in enumerate(train_loader, 0):\n        inputs = batch['image']\n        labels = batch['label']\n        inputs = inputs.to(device)\n        labels = labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n\n        outputs = torch.sigmoid(outputs)\n        loss = bce_loss(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        print(loss)\n    print(\"done epoch:\", epoch+1)\n    if (epoch + 1 == epochs):\n        print(\"saving model...\")\n        torch.save(model.state_dict(), os.path.join(model_dir, \"best_metric_model.pth\"))","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:22:19.416643Z","iopub.execute_input":"2022-07-20T01:22:19.417129Z","iopub.status.idle":"2022-07-20T01:22:38.767930Z","shell.execute_reply.started":"2022-07-20T01:22:19.417088Z","shell.execute_reply":"2022-07-20T01:22:38.766620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from monai.data import decollate_batch\n\nfrom monai.transforms import (Compose, Activations, AsDiscrete, Lambda)\n\npost_trans = Compose(\n    [AsDiscrete(threshold=0.0001)]\n)\n\npost_pred = Compose(\n    [AsDiscrete(threshold=0.5)]\n)\n\nlambd = Lambda(func=lambda x: x if x > 0 else 0) ","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:57.957442Z","iopub.status.idle":"2022-07-20T01:14:57.958162Z","shell.execute_reply.started":"2022-07-20T01:14:57.957911Z","shell.execute_reply":"2022-07-20T01:14:57.957934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %% Inspect predictions for first batch\nimport pandas as pd\n#inputs, labels = next(iter(test_loader))\nsample = next(iter(train_loader))\n#inputs, labels = first(test_loader)\ninputs = sample['image']\nlabels = sample['label']\ninputs = inputs.to(device)\nlabels = labels.numpy()\n#outputs = model(inputs).max(1).indices.detach().cpu().numpy()\noutputs = torch.sigmoid(model(inputs))\noutputs = outputs.detach().cpu()\noutputs = [post_pred(i) for i in decollate_batch(outputs)]\noutputs = [output.numpy() for output in outputs]\n\ncomparison = pd.DataFrame()\n#print(\"Batch accuracy: \", (labels==outputs).sum()/len(labels))\ncomparison[\"labels\"] = [label for label in labels]\n\ncomparison[\"outputs\"] = [output for output in outputs]\ncomparison","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:57.959455Z","iopub.status.idle":"2022-07-20T01:14:57.960178Z","shell.execute_reply.started":"2022-07-20T01:14:57.959929Z","shell.execute_reply":"2022-07-20T01:14:57.959953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision.transforms.functional as F","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:57.961467Z","iopub.status.idle":"2022-07-20T01:14:57.962203Z","shell.execute_reply.started":"2022-07-20T01:14:57.961943Z","shell.execute_reply":"2022-07-20T01:14:57.961967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %% Layerwise relevance propagation for VGG16\n# For other CNN architectures this code might become more complex\n# Source: https://git.tu-berlin.de/gmontavon/lrp-tutorial\n# http://iphome.hhi.de/samek/pdf/MonXAI19.pdf\n\ndef new_layer(layer, g):\n    \"\"\"Clone a layer and pass its parameters through the function g.\"\"\"\n    layer = copy.deepcopy(layer)\n    try: layer.weight = torch.nn.Parameter(g(layer.weight))\n    except AttributeError: pass\n    try: layer.bias = torch.nn.Parameter(g(layer.bias))\n    except AttributeError: pass\n    return layer\n\ndef dense_to_conv(layers):\n    \"\"\" Converts a dense layer to a conv layer \"\"\"\n    newlayers = []\n    for i,layer in enumerate(layers):\n        if isinstance(layer, nn.Linear):\n            newlayer = None\n            if i == 0:\n                m, n = 512, layer.weight.shape[0]\n                newlayer = nn.Conv2d(m,n,7)\n                newlayer.weight = nn.Parameter(layer.weight.reshape(n,m,7,7))\n            else:\n                m,n = layer.weight.shape[1],layer.weight.shape[0]\n                newlayer = nn.Conv2d(m,n,1)\n                newlayer.weight = nn.Parameter(layer.weight.reshape(n,m,1,1))\n            newlayer.bias = nn.Parameter(layer.bias)\n            newlayers += [newlayer]\n        else:\n            newlayers += [layer]\n    return newlayers\n\ndef get_linear_layer_indices(model):\n    offset = len(model.vgg16._modules['features']) + 1\n    indices = []\n    for i, layer in enumerate(model.vgg16._modules['classifier']): \n        if isinstance(layer, nn.Linear): \n            indices.append(i)\n    indices = [offset + val for val in indices]\n    return indices\n\ndef apply_lrp_on_vgg16(model, image):\n    image = torch.unsqueeze(image, 0)\n    # >>> Step 1: Extract layers\n    layers = list(model.vgg16._modules['features']) \\\n                + [model.vgg16._modules['avgpool']] \\\n                + dense_to_conv(list(model.vgg16._modules['classifier']))\n    linear_layer_indices = get_linear_layer_indices(model)\n    # >>> Step 2: Propagate image through layers and store activations\n    n_layers = len(layers)\n    activations = [image] + [None] * n_layers # list of activations\n        \n    for layer in range(n_layers):\n        if layer in linear_layer_indices:\n            if layer == 32:\n                activations[layer] = activations[layer].reshape((1, 512, 7, 7))\n        activation = layers[layer].forward(activations[layer])\n        if isinstance(layers[layer], torch.nn.modules.pooling.AdaptiveAvgPool2d):\n            activation = torch.flatten(activation, start_dim=1)\n        activations[layer+1] = activation\n        \n    # >>> Step 3: Replace last layer with one-hot-encoding\n    #outputs = model(inputs).max(1).indices.detach().cpu().numpy()\n\n    output_activation = activations[-1].detach().cpu().numpy()\n    print(\"output activation\", output_activation)\n    \n    outputs = torch.sigmoid(activations[-1])\n    outputs = outputs.detach().cpu()\n    outputs = [post_pred(i) for i in decollate_batch(outputs)]\n    outputs = [output.numpy() for output in outputs]\n    \n    max_activation = output_activation.max()\n    one_hot_output = [val if val == max_activation else 0 \n                        for val in output_activation[0]]\n    \n    out = []\n    for i in range(len((outputs)[0])):\n        if outputs[0][i]==1:\n            out.append(output_activation[0][i])\n        else:\n            out.append(0)\n    print(out)\n    \n    max_activation = output_activation.max()\n    one_hot_output = [val if val == max_activation else 0 \n                        for val in output_activation[0]]\n    #activations[-1] = torch.FloatTensor([one_hot_output]).to(device)\n    \n    activations[-1] = torch.FloatTensor([one_hot_output]).to(device)\n    \n    # >>> Step 4: Backpropagate relevance scores\n    relevances = [None] * n_layers + [activations[-1]]\n    # Iterate over the layers in reverse order\n    for layer in range(0, n_layers)[::-1]:\n        current = layers[layer]\n        # Treat max pooling layers as avg pooling\n        if isinstance(current, torch.nn.MaxPool2d):\n            layers[layer] = torch.nn.AvgPool2d(2)\n            current = layers[layer]\n        if isinstance(current, torch.nn.Conv2d) or \\\n           isinstance(current, torch.nn.AvgPool2d) or\\\n           isinstance(current, torch.nn.Linear):\n            activations[layer] = activations[layer].data.requires_grad_(True)\n            \n            # Apply variants of LRP depending on the depth\n            # see: https://link.springer.com/chapter/10.1007%2F978-3-030-28954-6_10\n            # Lower layers, LRP-gamma >> Favor positive contributions (activations)\n            if layer <= 16:       rho = lambda p: p + 0.25*p.clamp(min=0); incr = lambda z: z+1e-9\n            # Middle layers, LRP-epsilon >> Remove some noise / Only most salient factors survive\n            if 17 <= layer <= 30: rho = lambda p: p;                       incr = lambda z: z+1e-9+0.25*((z**2).mean()**.5).data\n            # Upper Layers, LRP-0 >> Basic rule\n            if layer >= 31:       rho = lambda p: p;                       incr = lambda z: z+1e-9\n            \n            # Transform weights of layer and execute forward pass\n            z = incr(new_layer(layers[layer],rho).forward(activations[layer]))\n            # Element-wise division between relevance of the next layer and z\n            s = (relevances[layer+1]/z).data                                     \n            # Calculate the gradient and multiply it by the activation\n            (z * s).sum().backward(); \n            c = activations[layer].grad       \n            # Assign new relevance values           \n            relevances[layer] = (activations[layer]*c).data                          \n        else:\n            relevances[layer] = relevances[layer+1]\n\n    # >>> Potential Step 5: Apply different propagation rule for pixels\n    return relevances[0]","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:57.963569Z","iopub.status.idle":"2022-07-20T01:14:57.964295Z","shell.execute_reply.started":"2022-07-20T01:14:57.964045Z","shell.execute_reply":"2022-07-20T01:14:57.964068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# %%\n# Calculate relevances for first image in this test batch\nimage_id = 26\nimage = torch.unsqueeze(inputs[image_id], 0)\n\nimage_relevances = apply_lrp_on_vgg16(model, inputs[image_id])\nimage_relevances = image_relevances.permute(0,2,3,1).detach().cpu().numpy()[0]\n'''\nimage_relevances = np.interp(image_relevances, (image_relevances.min(),\n                                                image_relevances.max()), \n                                                (0, 1))\n'''\n# Show relevances\n#pred_label = list(test_dataset.class_to_idx.keys())[list(test_dataset.class_to_idx.values()).index(labels[image_id])]\n\npred_label = labels[image_id]\n\nif outputs[image_id].all() == labels[image_id].all():\n    print(\"Groundtruth for this image: \", pred_label)\n    tensor_ir = torch.FloatTensor(image_relevances)\n    tensor_ir = tensor_ir.permute(2,0,1)\n    pos_ir = torch.nn.functional.relu(tensor_ir)\n    \n    post_ir = [post_trans(i) for i in decollate_batch(pos_ir)]\n    post_ir = torch.stack(post_ir)\n\n    scaled_ir = F.adjust_contrast(tensor_ir, 10)\n    \n    img_np = np.array(pos_ir)\n    print(np.mean(img_np))\n    # plot the pixel values\n    plt.hist(img_np.ravel(), bins=50, density=True)\n    plt.xlabel(\"pixel values\")\n    plt.ylabel(\"relative frequency\")\n    plt.title(\"distribution of pixels\")\n    \n    plt.figure(\"outputs\", (24,6))\n    # Plot images next to each other\n    plt.axis('off')\n    plt.subplot(1,3,1)\n    plt.title(\"original image relevance\")\n    plt.imshow(image_relevances[:,:,0], cmap=\"seismic\")\n    plt.subplot(1,3,2)\n    plt.title(\"asDiscrete image relevance\")\n    plt.imshow(post_ir[0,:,:])\n    plt.subplot(1,3,3)\n    plt.imshow(inputs[image_id].permute(1,2,0).detach().cpu().numpy())\n    plt.show()\nelse:\n    print(\"This image is not classified correctly.\")\n\n# %%","metadata":{"execution":{"iopub.status.busy":"2022-07-20T01:14:57.965625Z","iopub.status.idle":"2022-07-20T01:14:57.966336Z","shell.execute_reply.started":"2022-07-20T01:14:57.966085Z","shell.execute_reply":"2022-07-20T01:14:57.966109Z"},"trusted":true},"execution_count":null,"outputs":[]}]}