Browse Source

First version of Dataset and Dataloader implemented

Ruben Aguilo Schuurs 2 years ago
parent
commit
bc49cac40a
100 changed files with 11620 additions and 46 deletions
  1. 48 46
      main.py
  2. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I11767_masked_brain.nii.nii
  3. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I11879_masked_brain.nii.nii
  4. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I12061_masked_brain.nii.nii
  5. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I13509_masked_brain.nii.nii
  6. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I13807_masked_brain.nii.nii
  7. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I14166_masked_brain.nii.nii
  8. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I14808_masked_brain.nii.nii
  9. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I16169_masked_brain.nii.nii
  10. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I16238_masked_brain.nii.nii
  11. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I16740_masked_brain.nii.nii
  12. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I16828_masked_brain.nii.nii
  13. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I17363_masked_brain.nii.nii
  14. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I17415_masked_brain.nii.nii
  15. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I17585_masked_brain.nii.nii
  16. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I20080_masked_brain.nii.nii
  17. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I20332_masked_brain.nii.nii
  18. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I20506_masked_brain.nii.nii
  19. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I20771_masked_brain.nii.nii
  20. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I23153_masked_brain.nii.nii
  21. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I23431_masked_brain.nii.nii
  22. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I11771_masked_brain.nii.nii
  23. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I12151_masked_brain.nii.nii
  24. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I13447_masked_brain.nii.nii
  25. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I14114_masked_brain.nii.nii
  26. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I14783_masked_brain.nii.nii
  27. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I16867_masked_brain.nii.nii
  28. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I17508_masked_brain.nii.nii
  29. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I18931_masked_brain.nii.nii
  30. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I20550_masked_brain.nii.nii
  31. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I20726_masked_brain.nii.nii
  32. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I23089_masked_brain.nii.nii
  33. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I23667_masked_brain.nii.nii
  34. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I23798_masked_brain.nii.nii
  35. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I23901_masked_brain.nii.nii
  36. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I24641_masked_brain.nii.nii
  37. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I25715_masked_brain.nii.nii
  38. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I26030_masked_brain.nii.nii
  39. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I26940_masked_brain.nii.nii
  40. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I26947_masked_brain.nii.nii
  41. BIN
      original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I28549_masked_brain.nii.nii
  42. 0 0
      original_model/LP_ADNIMERGE.csv
  43. BIN
      original_model/SavedCNNModel
  44. BIN
      original_model/SavedCNNWeights
  45. BIN
      original_model/SavedRNNModel
  46. BIN
      original_model/SavedRNNWeights
  47. 11 0
      original_model/innvestigate/__init__.py
  48. 143 0
      original_model/innvestigate/analyzer/__init__.py
  49. 828 0
      original_model/innvestigate/analyzer/base.py
  50. 405 0
      original_model/innvestigate/analyzer/deeplift.py
  51. 204 0
      original_model/innvestigate/analyzer/deeptaylor.py
  52. 310 0
      original_model/innvestigate/analyzer/gradient_based.py
  53. 68 0
      original_model/innvestigate/analyzer/misc.py
  54. 272 0
      original_model/innvestigate/analyzer/pattern_based.py
  55. 1 0
      original_model/innvestigate/analyzer/relevance_based/__init__.py
  56. 869 0
      original_model/innvestigate/analyzer/relevance_based/relevance_analyzer.py
  57. 635 0
      original_model/innvestigate/analyzer/relevance_based/relevance_rule.py
  58. 102 0
      original_model/innvestigate/analyzer/relevance_based/utils.py
  59. 333 0
      original_model/innvestigate/analyzer/wrapper.py
  60. 0 0
      original_model/innvestigate/applications/__init__.py
  61. 297 0
      original_model/innvestigate/applications/imagenet.py
  62. 106 0
      original_model/innvestigate/applications/mnist.py
  63. 660 0
      original_model/innvestigate/layers.py
  64. 0 0
      original_model/innvestigate/tests/__init__.py
  65. 0 0
      original_model/innvestigate/tests/analyzer/__init__.py
  66. 231 0
      original_model/innvestigate/tests/analyzer/test_base.py
  67. 241 0
      original_model/innvestigate/tests/analyzer/test_deeplift.py
  68. 82 0
      original_model/innvestigate/tests/analyzer/test_deeptaylor.py
  69. 337 0
      original_model/innvestigate/tests/analyzer/test_gradient_based.py
  70. 51 0
      original_model/innvestigate/tests/analyzer/test_init.py
  71. 57 0
      original_model/innvestigate/tests/analyzer/test_misc.py
  72. 134 0
      original_model/innvestigate/tests/analyzer/test_pattern_based.py
  73. 242 0
      original_model/innvestigate/tests/analyzer/test_relevance_based.py
  74. 157 0
      original_model/innvestigate/tests/analyzer/test_wrapper.py
  75. 0 0
      original_model/innvestigate/tests/tools/__init__.py
  76. 382 0
      original_model/innvestigate/tests/tools/test_pattern.py
  77. 102 0
      original_model/innvestigate/tests/tools/test_perturbate.py
  78. 0 0
      original_model/innvestigate/tests/utils/__init__.py
  79. 0 0
      original_model/innvestigate/tests/utils/keras/__init__.py
  80. 80 0
      original_model/innvestigate/tests/utils/keras/test_graph.py
  81. 32 0
      original_model/innvestigate/tests/utils/test_visualizations.py
  82. 0 0
      original_model/innvestigate/tests/utils/tests/__init__.py
  83. 37 0
      original_model/innvestigate/tests/utils/tests/test_dryrun.py
  84. 64 0
      original_model/innvestigate/tests/utils/tests/test_layer.py
  85. 9 0
      original_model/innvestigate/tools/__init__.py
  86. 520 0
      original_model/innvestigate/tools/pattern.py
  87. 390 0
      original_model/innvestigate/tools/perturbate.py
  88. 185 0
      original_model/innvestigate/utils/__init__.py
  89. 77 0
      original_model/innvestigate/utils/keras/__init__.py
  90. 171 0
      original_model/innvestigate/utils/keras/backend.py
  91. 435 0
      original_model/innvestigate/utils/keras/checks.py
  92. 1152 0
      original_model/innvestigate/utils/keras/graph.py
  93. 0 0
      original_model/innvestigate/utils/tests/__init__.py
  94. 338 0
      original_model/innvestigate/utils/tests/dryrun.py
  95. 92 0
      original_model/innvestigate/utils/tests/layer.py
  96. 55 0
      original_model/innvestigate/utils/tests/networks/__init__.py
  97. 277 0
      original_model/innvestigate/utils/tests/networks/base.py
  98. 94 0
      original_model/innvestigate/utils/tests/networks/cifar10.py
  99. 210 0
      original_model/innvestigate/utils/tests/networks/imagenet.py
  100. 94 0
      original_model/innvestigate/utils/tests/networks/mnist.py

+ 48 - 46
main.py

@@ -1,14 +1,31 @@
 import torch
+import torchvision
 
 # FOR DATA
-from torch import nn
-from torch.utils.data import DataLoader, Dataset
+from utils.preprocess import prepare_datasets
+from utils.show_image import show_image
+from torch.utils.data import DataLoader
 from torchvision import datasets
+
+from torch import nn
 from torchvision.transforms import ToTensor
+
+# import nonechucks as nc     # Used to load data in pytorch even when images are corrupted / unavailable (skips them)
+
+# FOR IMAGE VISUALIZATION
+import nibabel as nib
+
+# GENERAL PURPOSE
 import os
 import pandas as pd
-from torchvision.io import read_image
-import nonechucks as nc     # Used to load data in pytorch even when images are corrupted / unavailable (skips them)
+import numpy as np
+import matplotlib.pyplot as plt
+import glob
+
+
+
+print("--- RUNNING ---")
+print("Pytorch Version: " + torch. __version__)
 
 
 # MAYBE??
@@ -21,11 +38,10 @@ os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
 os.environ["CUDA_VISIBLE_DEVICES"] = "0"  # use id from $ nvidia-smi
 '''
 
-
-
-
 # LOADING DATA
 # data & training properties:
+val_split = 0.2     # % of val and test, rest will be train
+seed = 12       # TODO Randomize seed
 '''
 target_rows = 91
 target_cols = 109
@@ -43,59 +59,45 @@ optimizer = Adam(lr=1e-5)
 final_layer_size = 5
 '''
 
-# Defines Dataset class, which is instantiated to create the dataset with the required images
-class CustomImageDataset(Dataset):
-    def __init__(self, annotations_file, img_dir, transform=None, target_transform=None):
-        self.img_labels = pd.read_csv(annotations_file)
-        self.img_dir = img_dir
-        self.transform = transform
-        self.target_transform = target_transform
-
-    def __len__(self):
-        return len(self.img_labels)
-
-    def __getitem__(self, idx):
-        img_path = os.path.join(self.img_dir, self.img_labels.iloc[idx, 0])
-        image = read_image(img_path)
-        label = self.img_labels.iloc[idx, 1]
-        if self.transform:
-            image = self.transform(image)
-        if self.target_transform:
-            label = self.target_transform(label)
-        return image, label
-
 
 # Might have to replace datapaths or separate between training and testing
-model_filepath = '//data/data_wnx1/rschuurs/pytorchCNN'
+model_filepath = '//data/data_wnx1/rschuurs/Pytorch_CNN-RNN'
 mri_datapath = './ADNI_volumes_customtemplate_float32/'
-anotationsFileDataPath = './LP_ADNIMERGE.csv'
+annotations_datapath = './LP_ADNIMERGE.csv'
 
-training_data = CustomImageDataset(anotationsFileDataPath, mri_datapath)
-test_data = CustomImageDataset(anotationsFileDataPath, mri_datapath)
+# annotations_file = pd.read_csv(annotations_datapath)    # DataFrame
 
-# Skips corrupted or non-present files in CSV
-safeTrainingData = nc.SafeDataset(training_data)
-safeTestData = nc.SafeDataset(test_data)
-
-print("Initial Training data length: " + str(training_data.__len__()))
-print("Safe Training data length: " + str(safeTrainingData.__len__()))
-print("Test data length: " + str(training_data.__len__()))
+# show_image(17508)
 
+# TODO: Datasets include multiple labels, such as medical info
+training_data, val_data, test_data = prepare_datasets(mri_datapath, val_split, seed)
 batch_size = 64
 
-# Create data loaders.
-train_dataloader = DataLoader(training_data, batch_size=batch_size)
-test_dataloader = DataLoader(test_data, batch_size=batch_size)
+# Create data loaders
+train_dataloader = DataLoader(training_data, batch_size=batch_size, shuffle=True)
+test_dataloader = DataLoader(test_data, batch_size=batch_size, shuffle=True)
+val_dataloader = DataLoader(val_data, batch_size=batch_size, shuffle=True)
 
-for X, y in test_dataloader:
+for X, y in train_dataloader:
     print(f"Shape of X [N, C, H, W]: {X.shape}")
     print(f"Shape of y: {y.shape} {y.dtype}")
     break
 
 
+# Display 10 images and labels.
+x = 0
+while x < 10:
+    train_features, train_labels = next(iter(train_dataloader))
+    print(f"Feature batch shape: {train_features.size()}")
+    img = train_features[0].squeeze()
+    image = img[:, :, 40]
+    label = train_labels[0]
+    plt.imshow(image, cmap="gray")
+    plt.show()
+    print(f"Label: {label}")
+    x = x+1
 
-
-
+print("--- END ---")
 
 # EXTRA
 
@@ -122,4 +124,4 @@ def evaluate_net (seed):
     train_data, val_data, test_data,rnn_HdataT1,rnn_HdataT2,rnn_HdataT3,rnn_AdataT1,rnn_AdataT2,rnn_AdataT3, test_mri_nonorm = data_loader.get_train_val_test(val_split, mri_datapath)
 
     print('Length Val Data[0]: ',len(val_data[0]))
-'''
+'''

BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I11767_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I11879_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I12061_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I13509_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I13807_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I14166_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I14808_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I16169_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I16238_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I16740_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I16828_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I17363_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I17415_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I17585_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I20080_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I20332_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I20506_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I20771_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I23153_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableAD__I23431_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I11771_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I12151_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I13447_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I14114_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I14783_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I16867_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I17508_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I18931_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I20550_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I20726_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I23089_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I23667_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I23798_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I23901_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I24641_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I25715_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I26030_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I26940_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I26947_masked_brain.nii.nii


BIN
original_model/ADNI_volumes_customtemplate_float32/Inf_NaN_stableNL__I28549_masked_brain.nii.nii


File diff suppressed because it is too large
+ 0 - 0
original_model/LP_ADNIMERGE.csv


BIN
original_model/SavedCNNModel


BIN
original_model/SavedCNNWeights


BIN
original_model/SavedRNNModel


BIN
original_model/SavedRNNWeights


+ 11 - 0
original_model/innvestigate/__init__.py

@@ -0,0 +1,11 @@
+
+from . import analyzer
+from .analyzer import create_analyzer
+from .analyzer import NotAnalyzeableModelException
+
+# Disable pyflaks warnings:
+assert analyzer
+assert create_analyzer
+assert NotAnalyzeableModelException
+
+__version__ = '1.0.9'

+ 143 - 0
original_model/innvestigate/analyzer/__init__.py

@@ -0,0 +1,143 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+from .base import NotAnalyzeableModelException
+from .deeplift import DeepLIFT
+from .deeplift import DeepLIFTWrapper
+from .gradient_based import BaselineGradient
+from .gradient_based import Gradient
+from .gradient_based import InputTimesGradient
+from .gradient_based import GuidedBackprop
+from .gradient_based import Deconvnet
+from .gradient_based import IntegratedGradients
+from .gradient_based import SmoothGrad
+from .misc import Input
+from .misc import Random
+from .pattern_based import PatternNet
+from .pattern_based import PatternAttribution
+from .relevance_based.relevance_analyzer import BaselineLRPZ
+from .relevance_based.relevance_analyzer import LRP
+from .relevance_based.relevance_analyzer import LRPZ
+from .relevance_based.relevance_analyzer import LRPZIgnoreBias
+from .relevance_based.relevance_analyzer import LRPZPlus
+from .relevance_based.relevance_analyzer import LRPZPlusFast
+from .relevance_based.relevance_analyzer import LRPEpsilon
+from .relevance_based.relevance_analyzer import LRPEpsilonIgnoreBias
+from .relevance_based.relevance_analyzer import LRPWSquare
+from .relevance_based.relevance_analyzer import LRPFlat
+from .relevance_based.relevance_analyzer import LRPAlphaBeta
+from .relevance_based.relevance_analyzer import LRPAlpha2Beta1
+from .relevance_based.relevance_analyzer import LRPAlpha2Beta1IgnoreBias
+from .relevance_based.relevance_analyzer import LRPAlpha1Beta0
+from .relevance_based.relevance_analyzer import LRPAlpha1Beta0IgnoreBias
+from .relevance_based.relevance_analyzer import LRPSequentialPresetA
+from .relevance_based.relevance_analyzer import LRPSequentialPresetB
+from .relevance_based.relevance_analyzer import LRPSequentialPresetAFlat
+from .relevance_based.relevance_analyzer import LRPSequentialPresetBFlat
+from .relevance_based.relevance_analyzer import LRPSequentialPresetBFlatUntilIdx
+from .deeptaylor import DeepTaylor
+from .deeptaylor import BoundedDeepTaylor
+from .wrapper import WrapperBase
+from .wrapper import AugmentReduceBase
+from .wrapper import GaussianSmoother
+from .wrapper import PathIntegrator
+
+
+# Disable pyflaks warnings:
+assert NotAnalyzeableModelException
+assert DeepLIFT
+assert BaselineLRPZ
+assert WrapperBase
+assert AugmentReduceBase
+assert GaussianSmoother
+assert PathIntegrator
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+analyzers = {
+    # Utility.
+    "input": Input,
+    "random": Random,
+
+    # Gradient based
+    "gradient": Gradient,
+    "gradient.baseline": BaselineGradient,
+    "input_t_gradient": InputTimesGradient,
+    "deconvnet": Deconvnet,
+    "guided_backprop": GuidedBackprop,
+    "integrated_gradients": IntegratedGradients,
+    "smoothgrad": SmoothGrad,
+
+    # Relevance based
+    "lrp": LRP,
+    "lrp.z": LRPZ,
+    "lrp.z_IB": LRPZIgnoreBias,
+
+    "lrp.epsilon": LRPEpsilon,
+    "lrp.epsilon_IB": LRPEpsilonIgnoreBias,
+
+    "lrp.w_square": LRPWSquare,
+    "lrp.flat": LRPFlat,
+
+    "lrp.alpha_beta": LRPAlphaBeta,
+
+    "lrp.alpha_2_beta_1": LRPAlpha2Beta1,
+    "lrp.alpha_2_beta_1_IB": LRPAlpha2Beta1IgnoreBias,
+    "lrp.alpha_1_beta_0": LRPAlpha1Beta0,
+    "lrp.alpha_1_beta_0_IB": LRPAlpha1Beta0IgnoreBias,
+    "lrp.z_plus": LRPZPlus,
+    "lrp.z_plus_fast": LRPZPlusFast,
+
+    "lrp.sequential_preset_a": LRPSequentialPresetA,
+    "lrp.sequential_preset_b": LRPSequentialPresetB,
+    "lrp.sequential_preset_a_flat": LRPSequentialPresetAFlat,
+    "lrp.sequential_preset_b_flat": LRPSequentialPresetBFlat,
+    "lrp.sequential_preset_b_flat_until_idx": LRPSequentialPresetBFlatUntilIdx,
+
+
+    # Deep Taylor
+    "deep_taylor": DeepTaylor,
+    "deep_taylor.bounded": BoundedDeepTaylor,
+
+    # DeepLIFT
+    #"deep_lift": DeepLIFT,
+    "deep_lift.wrapper": DeepLIFTWrapper,
+
+    # Pattern based
+    "pattern.net": PatternNet,
+    "pattern.attribution": PatternAttribution,
+}
+
+
+def create_analyzer(name, model, **kwargs):
+    """Instantiates the analyzer with the name 'name'
+
+    This convenience function takes an analyzer name
+    creates the respective analyzer.
+
+    Alternatively analyzers can be created directly by
+    instantiating the respective classes.
+
+    :param name: Name of the analyzer.
+    :param model: The model to analyze, passed to the analyzer's __init__.
+    :param kwargs: Additional parameters for the analyzer's .
+    :return: An instance of the chosen analyzer.
+    :raise KeyError: If there is no analyzer with the passed name.
+    """
+    try:
+        analyzer_class = analyzers[name]
+    except KeyError:
+        raise KeyError(
+            "No analyzer with the name '%s' could be found."
+            " All possible names are: %s" % (name, list(analyzers.keys())))
+    return analyzer_class(model, **kwargs)

+ 828 - 0
original_model/innvestigate/analyzer/base.py

@@ -0,0 +1,828 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import zip
+import six
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.backend as K
+import keras.layers
+import keras.models
+import numpy as np
+import warnings
+
+
+from .. import layers as ilayers
+from .. import utils as iutils
+from ..utils.keras import checks as kchecks
+from ..utils.keras import graph as kgraph
+
+
+__all__ = [
+    "NotAnalyzeableModelException",
+    "AnalyzerBase",
+
+    "TrainerMixin",
+    "OneEpochTrainerMixin",
+
+    "AnalyzerNetworkBase",
+    "ReverseAnalyzerBase"
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class NotAnalyzeableModelException(Exception):
+    """Indicates that the model cannot be analyzed by an analyzer."""
+    pass
+
+
+class AnalyzerBase(object):
+    """ The basic interface of an iNNvestigate analyzer.
+
+    This class defines the basic interface for analyzers:
+
+    >>> model = create_keras_model()
+    >>> a = Analyzer(model)
+    >>> a.fit(X_train)  # If analyzer needs training.
+    >>> analysis = a.analyze(X_test)
+    >>>
+    >>> state = a.save()
+    >>> a_new = A.load(*state)
+    >>> analysis = a_new.analyze(X_test)
+
+    :param model: A Keras model.
+    :param disable_model_checks: Do not execute model checks that enforce
+      compatibility of analyzer and model.
+
+    .. note:: To develop a new analyzer derive from
+      :class:`AnalyzerNetworkBase`.
+    """
+
+    def __init__(self, model, disable_model_checks=False):
+        self._model = model
+        self._disable_model_checks = disable_model_checks
+
+        self._do_model_checks()
+
+    def _add_model_check(self, check, message, check_type="exception"):
+        if getattr(self, "_model_check_done", False):
+            raise Exception("Cannot add model check anymore."
+                            " Check was already performed.")
+
+        if not hasattr(self, "_model_checks"):
+            self._model_checks = []
+
+        check_instance = {
+            "check": check,
+            "message": message,
+            "type": check_type,
+        }
+        self._model_checks.append(check_instance)
+
+    def _do_model_checks(self):
+        model_checks = getattr(self, "_model_checks", [])
+
+        if not self._disable_model_checks and len(model_checks) > 0:
+            check = [x["check"] for x in model_checks]
+            types = [x["type"] for x in model_checks]
+            messages = [x["message"] for x in model_checks]
+
+            checked = kgraph.model_contains(self._model, check)
+            tmp = zip(iutils.to_list(checked), messages, types)
+
+            for checked_layers, message, check_type in tmp:
+                if len(checked_layers) > 0:
+                    tmp_message = ("%s\nCheck triggerd by layers: %s" %
+                                   (message, checked_layers))
+
+                    if check_type == "exception":
+                        raise NotAnalyzeableModelException(tmp_message)
+                    elif check_type == "warning":
+                        # TODO(albermax) only the first warning will be shown
+                        warnings.warn(tmp_message)
+                    else:
+                        raise NotImplementedError()
+
+        self._model_check_done = True
+
+    def fit(self, *args, **kwargs):
+        """
+        Stub that eats arguments. If an analyzer needs training
+        include :class:`TrainerMixin`.
+
+        :param disable_no_training_warning: Do not warn if this function is
+          called despite no training is needed.
+        """
+        disable_no_training_warning = kwargs.pop("disable_no_training_warning",
+                                                 False)
+        if not disable_no_training_warning:
+            # issue warning if not training is foreseen,
+            # but is fit is still called.
+            warnings.warn("This analyzer does not need to be trained."
+                          " Still fit() is called.", RuntimeWarning)
+
+    def fit_generator(self, *args, **kwargs):
+        """
+        Stub that eats arguments. If an analyzer needs training
+        include :class:`TrainerMixin`.
+
+        :param disable_no_training_warning: Do not warn if this function is
+          called despite no training is needed.
+        """
+        disable_no_training_warning = kwargs.pop("disable_no_training_warning",
+                                                 False)
+        if not disable_no_training_warning:
+            # issue warning if not training is foreseen,
+            # but is fit is still called.
+            warnings.warn("This analyzer does not need to be trained."
+                          " Still fit_generator() is called.", RuntimeWarning)
+
+    def analyze(self, X):
+        """
+        Analyze the behavior of model on input `X`.
+
+        :param X: Input as expected by model.
+        """
+        raise NotImplementedError()
+
+    def _get_state(self):
+        state = {
+            "model_json": self._model.to_json(),
+            "model_weights": self._model.get_weights(),
+            "disable_model_checks": self._disable_model_checks,
+        }
+        return state
+
+    def save(self):
+        """
+        Save state of analyzer, can be passed to :func:`Analyzer.load`
+        to resemble the analyzer.
+
+        :return: The class name and the state.
+        """
+        state = self._get_state()
+        class_name = self.__class__.__name__
+        return class_name, state
+
+    def save_npz(self, fname):
+        """
+        Save state of analyzer, can be passed to :func:`Analyzer.load_npz`
+        to resemble the analyzer.
+
+        :param fname: The file's name.
+        """
+        class_name, state = self.save()
+        np.savez(fname, **{"class_name": class_name,
+                           "state": state})
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        model_json = state.pop("model_json")
+        model_weights = state.pop("model_weights")
+        disable_model_checks = state.pop("disable_model_checks")
+        assert len(state) == 0
+
+        model = keras.models.model_from_json(model_json)
+        model.set_weights(model_weights)
+        return {"model": model,
+                "disable_model_checks": disable_model_checks}
+
+    @staticmethod
+    def load(class_name, state):
+        """
+        Resembles an analyzer from the state created by
+        :func:`analyzer.save()`.
+
+        :param class_name: The analyzer's class name.
+        :param state: The analyzer's state.
+        """
+        # Todo:do in a smarter way!
+        import innvestigate.analyzer
+        clazz = getattr(innvestigate.analyzer, class_name)
+
+        kwargs = clazz._state_to_kwargs(state)
+        return clazz(**kwargs)
+
+    @staticmethod
+    def load_npz(fname):
+        """
+        Resembles an analyzer from the file created by
+        :func:`analyzer.save_npz()`.
+
+        :param fname: The file's name.
+        """
+        f = np.load(fname)
+
+        class_name = f["class_name"].item()
+        state = f["state"].item()
+        return AnalyzerBase.load(class_name, state)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class TrainerMixin(object):
+    """Mixin for analyzer that adapt to data.
+
+    This convenience interface exposes a Keras like training routing
+    to the user.
+    """
+
+    # todo: extend with Y
+    def fit(self,
+            X=None,
+            batch_size=32,
+            **kwargs):
+        """
+        Takes the same parameters as Keras's :func:`model.fit` function.
+        """
+        generator = iutils.BatchSequence(X, batch_size)
+        return self._fit_generator(generator,
+                                  **kwargs)
+
+    def fit_generator(self, *args, **kwargs):
+        """
+        Takes the same parameters as Keras's :func:`model.fit_generator`
+        function.
+        """
+        return self._fit_generator(*args, **kwargs)
+
+    def _fit_generator(self,
+                       generator,
+                       steps_per_epoch=None,
+                       epochs=1,
+                       max_queue_size=10,
+                       workers=1,
+                       use_multiprocessing=False,
+                       verbose=0,
+                       disable_no_training_warning=None):
+        raise NotImplementedError()
+
+
+class OneEpochTrainerMixin(TrainerMixin):
+    """Exposes the same interface and functionality as :class:`TrainerMixin`
+    except that the training is limited to one epoch.
+    """
+
+    def fit(self, *args, **kwargs):
+        """
+        Same interface as :func:`fit` of :class:`TrainerMixin` except that
+        the parameter epoch is fixed to 1.
+        """
+        return super(OneEpochTrainerMixin, self).fit(*args, epochs=1, **kwargs)
+
+    def fit_generator(self, *args, **kwargs):
+        """
+        Same interface as :func:`fit_generator` of :class:`TrainerMixin` except that
+        the parameter epoch is fixed to 1.
+        """
+        steps = kwargs.pop("steps", None)
+        return super(OneEpochTrainerMixin, self).fit_generator(
+            *args,
+            steps_per_epoch=steps,
+            epochs=1,
+            **kwargs)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class AnalyzerNetworkBase(AnalyzerBase):
+    """Convenience interface for analyzers.
+
+    This class provides helpful functionality to create analyzer's.
+    Basically it:
+
+    * takes the input model and adds a layer that selects
+      the desired output neuron to analyze.
+    * passes the new model to :func:`_create_analysis` which should
+      return the analysis as Keras tensors.
+    * compiles the function and serves the output to :func:`analyze` calls.
+    * allows :func:`_create_analysis` to return tensors
+      that are intercept for debugging purposes.
+
+    :param neuron_selection_mode: How to select the neuron to analyze.
+      Possible values are 'max_activation', 'index' for the neuron
+      (expects indices at :func:`analyze` calls), 'all' take all neurons.
+    :param allow_lambda_layers: Allow the model to contain lambda layers.
+    """
+
+    def __init__(self, model,
+                 neuron_selection_mode="max_activation",
+                 allow_lambda_layers=False,
+                 **kwargs):
+        if neuron_selection_mode not in ["max_activation", "index", "all"]:
+            raise ValueError("neuron_selection parameter is not valid.")
+        self._neuron_selection_mode = neuron_selection_mode
+
+        self._allow_lambda_layers = allow_lambda_layers
+        self._add_model_check(
+            lambda layer: (not self._allow_lambda_layers and
+                           isinstance(layer, keras.layers.core.Lambda)),
+            ("Lamda layers are not allowed. "
+             "To force use set allow_lambda_layers parameter."),
+            check_type="exception",
+        )
+
+        self._special_helper_layers = []
+
+        super(AnalyzerNetworkBase, self).__init__(model, **kwargs)
+
+    def _add_model_softmax_check(self):
+        """
+        Adds check that prevents models from containing a softmax.
+        """
+        self._add_model_check(
+            lambda layer: kchecks.contains_activation(
+                layer, activation="softmax"),
+            "This analysis method does not support softmax layers.",
+            check_type="exception",
+        )
+
+    def _prepare_model(self, model):
+        """
+        Prepares the model to analyze before it gets actually analyzed.
+
+        This class adds the code to select a specific output neuron.
+        """
+        neuron_selection_mode = self._neuron_selection_mode
+        model_inputs = model.inputs
+
+        model_output = model.outputs
+        if len(model_output) > 1:
+            raise ValueError("Only models with one output tensor are allowed.")
+        analysis_inputs = []
+        stop_analysis_at_tensors = []
+
+        # Flatten to form (batch_size, other_dimensions):
+        if K.ndim(model_output[0]) > 2:
+            model_output = keras.layers.Flatten()(model_output)
+
+        if neuron_selection_mode == "max_activation":
+            l = ilayers.Max(name="iNNvestigate_max")
+            model_output = l(model_output)
+            self._special_helper_layers.append(l)
+        elif neuron_selection_mode == "index":
+            neuron_indexing = keras.layers.Input(
+                batch_shape=[None, None], dtype=np.int32,
+                name='iNNvestigate_neuron_indexing')
+            self._special_helper_layers.append(
+                neuron_indexing._keras_history[0])
+            analysis_inputs.append(neuron_indexing)
+            # The indexing tensor should not be analyzed.
+            stop_analysis_at_tensors.append(neuron_indexing)
+
+            l = ilayers.GatherND(name="iNNvestigate_gather_nd")
+            model_output = l(model_output+[neuron_indexing])
+            self._special_helper_layers.append(l)
+        elif neuron_selection_mode == "all":
+            pass
+        else:
+            raise NotImplementedError()
+        
+        model = keras.models.Model(inputs=model_inputs+analysis_inputs,
+                                   outputs=model_output)
+        return model, analysis_inputs, stop_analysis_at_tensors
+
+    def create_analyzer_model(self):
+        """
+        Creates the analyze functionality. If not called beforehand
+        it will be called by :func:`analyze`.
+        """
+        model_inputs = self._model.inputs
+        tmp = self._prepare_model(self._model)
+        model, analysis_inputs, stop_analysis_at_tensors = tmp
+        self._analysis_inputs = analysis_inputs
+        self._prepared_model = model
+
+        tmp = self._create_analysis(
+            model, stop_analysis_at_tensors=stop_analysis_at_tensors)
+        if isinstance(tmp, tuple):
+            if len(tmp) == 3:
+                analysis_outputs, debug_outputs, constant_inputs = tmp
+            elif len(tmp) == 2:
+                analysis_outputs, debug_outputs = tmp
+                constant_inputs = list()
+            elif len(tmp) == 1:
+                analysis_outputs = iutils.to_list(tmp[0])
+                constant_inputs, debug_outputs = list(), list()
+            else:
+                raise Exception("Unexpected output from _create_analysis.")
+        else:
+            analysis_outputs = tmp
+            constant_inputs, debug_outputs = list(), list()
+
+        analysis_outputs = iutils.to_list(analysis_outputs)
+        debug_outputs = iutils.to_list(debug_outputs)
+        constant_inputs = iutils.to_list(constant_inputs)
+
+        self._n_data_input = len(model_inputs)
+        self._n_constant_input = len(constant_inputs)
+        self._n_data_output = len(analysis_outputs)
+        self._n_debug_output = len(debug_outputs)
+        self._analyzer_model = keras.models.Model(
+            inputs=model_inputs+analysis_inputs+constant_inputs,
+            outputs=analysis_outputs+debug_outputs)
+
+    def _create_analysis(self, model, stop_analysis_at_tensors=[]):
+        """
+        Interface that needs to be implemented by a derived class.
+
+        This function is expected to create a Keras graph that creates
+        a custom analysis for the model inputs given the model outputs.
+
+        :param model: Target of analysis.
+        :param stop_analysis_at_tensors: A list of tensors where to stop the
+          analysis. Similar to stop_gradient arguments when computing the
+          gradient of a graph.
+        :return: Either one-, two- or three-tuple of lists of tensors.
+          * The first list of tensors represents the analysis for each
+            model input tensor. Tensors present in stop_analysis_at_tensors
+            should be omitted.
+          * The second list, if present, is a list of debug tensors that will
+            be passed to :func:`_handle_debug_output` after the analysis
+            is executed.
+          * The third list, if present, is a list of constant input tensors
+            added to the analysis model.
+        """
+        raise NotImplementedError()
+
+    def _handle_debug_output(self, debug_values):
+        raise NotImplementedError()
+
+    def analyze(self, X, neuron_selection=None):
+        """
+        Same interface as :class:`Analyzer` besides
+
+        :param neuron_selection: If neuron_selection_mode is 'index' this
+          should be an integer with the index for the chosen neuron.
+        """
+        if not hasattr(self, "_analyzer_model"):
+            self.create_analyzer_model()
+
+        X = iutils.to_list(X)
+
+        if(neuron_selection is not None and
+           self._neuron_selection_mode != "index"):
+            raise ValueError("Only neuron_selection_mode 'index' expects "
+                             "the neuron_selection parameter.")
+        if(neuron_selection is None and
+           self._neuron_selection_mode == "index"):
+            raise ValueError("neuron_selection_mode 'index' expects "
+                             "the neuron_selection parameter.")
+
+        if self._neuron_selection_mode == "index":
+            neuron_selection = np.asarray(neuron_selection).flatten()
+            if neuron_selection.size == 1:
+                neuron_selection = np.repeat(neuron_selection, len(X[0]))
+
+            # Add first axis indices for gather_nd
+            neuron_selection = np.hstack(
+                (np.arange(len(neuron_selection)).reshape((-1, 1)),
+                 neuron_selection.reshape((-1, 1)))
+            )
+            ret = self._analyzer_model.predict_on_batch(X+[neuron_selection])
+        else:
+            ret = self._analyzer_model.predict_on_batch(X)
+
+        if self._n_debug_output > 0:
+            self._handle_debug_output(ret[-self._n_debug_output:])
+            ret = ret[:-self._n_debug_output]
+
+        if isinstance(ret, list) and len(ret) == 1:
+            ret = ret[0]
+        return ret
+
+    def _get_state(self):
+        state = super(AnalyzerNetworkBase, self)._get_state()
+        state.update({"neuron_selection_mode": self._neuron_selection_mode})
+        state.update({"allow_lambda_layers": self._allow_lambda_layers})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        neuron_selection_mode = state.pop("neuron_selection_mode")
+        allow_lambda_layers = state.pop("allow_lambda_layers")
+        kwargs = super(AnalyzerNetworkBase, clazz)._state_to_kwargs(state)
+        kwargs.update({
+            "neuron_selection_mode": neuron_selection_mode,
+            "allow_lambda_layers": allow_lambda_layers
+        })
+        return kwargs
+
+
+class ReverseAnalyzerBase(AnalyzerNetworkBase):
+    """Convenience class for analyzers that revert the model's structure.
+
+    This class contains many helper functions around the graph
+    reverse function :func:`innvestigate.utils.keras.graph.reverse_model`.
+
+    The deriving classes should specify how the graph should be reverted
+    by implementing the following functions:
+
+    * :func:`_reverse_mapping(layer)` given a layer this function
+      returns a reverse mapping for the layer as specified in
+      :func:`innvestigate.utils.keras.graph.reverse_model` or None.
+
+      This function can be implemented, but it is encouraged to
+      implement a default mapping and add additional changes with
+      the function :func:`_add_conditional_reverse_mapping` (see below).
+
+      The default behavior is finding a conditional mapping (see below),
+      if none is found, :func:`_default_reverse_mapping` is applied.
+    * :func:`_default_reverse_mapping` defines the default
+      reverse mapping.
+    * :func:`_head_mapping` defines how the outputs of the model
+      should be instantiated before the are passed to the reversed
+      network.
+
+    Furthermore other parameters of the function
+    :func:`innvestigate.utils.keras.graph.reverse_model` can
+    be changed by setting the according parameters of the
+    init function:
+
+    :param reverse_verbose: Print information on the reverse process.
+    :param reverse_clip_values: Clip the values that are passed along
+      the reverted network. Expects tuple (min, max).
+    :param reverse_project_bottleneck_layers: Project the value range
+      of bottleneck tensors in the reverse network into another range.
+    :param reverse_check_min_max_values: Print the min/max values
+      observed in each tensor along the reverse network whenever
+      :func:`analyze` is called.
+    :param reverse_check_finite: Check if values passed along the
+      reverse network are finite.
+    :param reverse_keep_tensors: Keeps the tensors created in the
+      backward pass and stores them in the attribute
+      :attr:`_reversed_tensors`.
+    :param reverse_reapply_on_copied_layers: See
+      :func:`innvestigate.utils.keras.graph.reverse_model`.
+    """
+
+    def __init__(self,
+                 model,
+                 reverse_verbose=False,
+                 reverse_clip_values=False,
+                 reverse_project_bottleneck_layers=False,
+                 reverse_check_min_max_values=False,
+                 reverse_check_finite=False,
+                 reverse_keep_tensors=False,
+                 reverse_reapply_on_copied_layers=False,
+                 **kwargs):
+        self._reverse_verbose = reverse_verbose
+        self._reverse_clip_values = reverse_clip_values
+        self._reverse_project_bottleneck_layers = (
+            reverse_project_bottleneck_layers)
+        self._reverse_check_min_max_values = reverse_check_min_max_values
+        self._reverse_check_finite = reverse_check_finite
+        self._reverse_keep_tensors = reverse_keep_tensors
+        self._reverse_reapply_on_copied_layers = (
+            reverse_reapply_on_copied_layers)
+        super(ReverseAnalyzerBase, self).__init__(model, **kwargs)
+
+    def _gradient_reverse_mapping(self, Xs, Ys, reversed_Ys, reverse_state):
+        mask = [x not in reverse_state["stop_mapping_at_tensors"] for x in Xs]
+        return ilayers.GradientWRT(len(Xs), mask=mask)(Xs+Ys+reversed_Ys)
+
+    def _reverse_mapping(self, layer):
+        """
+        This function should return a reverse mapping for the passed layer.
+
+        If this function returns None, :func:`_default_reverse_mapping`
+        is applied.
+
+        :param layer: The layer for which a mapping should be returned.
+        :return: The mapping can be of the following forms:
+          * A function of form (A) f(Xs, Ys, reversed_Ys, reverse_state)
+            that maps reversed_Ys to reversed_Xs (which should contain
+            tensors of the same shape and type).
+          * A function of form f(B) f(layer, reverse_state) that returns
+            a function of form (A).
+          * A :class:`ReverseMappingBase` subclass.
+        """
+        if layer in self._special_helper_layers:
+            # Special layers added by AnalyzerNetworkBase
+            # that should not be exposed to user.
+            return self._gradient_reverse_mapping
+
+        return self._apply_conditional_reverse_mappings(layer)
+
+    def _add_conditional_reverse_mapping(
+            self, condition, mapping, priority=-1, name=None):
+        """
+        This function should return a reverse mapping for the passed layer.
+
+        If this function returns None, :func:`_default_reverse_mapping`
+        is applied.
+
+        :param condition: Condition when this mapping should be applied.
+          Form: f(layer) -> bool
+        :param mapping: The mapping can be of the following forms:
+          * A function of form (A) f(Xs, Ys, reversed_Ys, reverse_state)
+            that maps reversed_Ys to reversed_Xs (which should contain
+            tensors of the same shape and type).
+          * A function of form f(B) f(layer, reverse_state) that returns
+            a function of form (A).
+          * A :class:`ReverseMappingBase` subclass.
+        :param priority: The higher the earlier the condition gets
+          evaluated.
+        :param name: An identifying name.
+        """
+        if getattr(self, "_reverse_mapping_applied", False):
+            raise Exception("Cannot add conditional mapping "
+                            "after first application.")
+
+        if not hasattr(self, "_conditional_reverse_mappings"):
+            self._conditional_reverse_mappings = {}
+
+        if priority not in self._conditional_reverse_mappings:
+            self._conditional_reverse_mappings[priority] = []
+
+        tmp = {"condition": condition, "mapping": mapping, "name": name}
+        self._conditional_reverse_mappings[priority].append(tmp)
+
+    def _apply_conditional_reverse_mappings(self, layer):
+        mappings = getattr(self, "_conditional_reverse_mappings", {})
+        self._reverse_mapping_applied = True
+
+        # Search for mapping. First consider ones with highest priority,
+        # inside priority in order of adding.
+        sorted_keys = sorted(mappings.keys())[::-1]
+        for key in sorted_keys:
+            for mapping in mappings[key]:
+                if mapping["condition"](layer):
+                    return mapping["mapping"]
+
+        return None
+
+    def _default_reverse_mapping(self, Xs, Ys, reversed_Ys, reverse_state):
+        """
+        Fallback function to map reversed_Ys to reversed_Xs
+        (which should contain tensors of the same shape and type).
+        """
+        return self._gradient_reverse_mapping(
+            Xs, Ys, reversed_Ys, reverse_state)
+
+    def _head_mapping(self, X):
+        """
+        Map output tensors to new values before passing
+        them into the reverted network.
+        """
+        return X
+
+    def _postprocess_analysis(self, X):
+        return X
+
+    def _reverse_model(self,
+                       model,
+                       stop_analysis_at_tensors=[],
+                       return_all_reversed_tensors=False):
+        return kgraph.reverse_model(
+            model,
+            reverse_mappings=self._reverse_mapping,
+            default_reverse_mapping=self._default_reverse_mapping,
+            head_mapping=self._head_mapping,
+            stop_mapping_at_tensors=stop_analysis_at_tensors,
+            verbose=self._reverse_verbose,
+            clip_all_reversed_tensors=self._reverse_clip_values,
+            project_bottleneck_tensors=self._reverse_project_bottleneck_layers,
+            return_all_reversed_tensors=return_all_reversed_tensors)
+
+    def _create_analysis(self, model, stop_analysis_at_tensors=[]):
+        return_all_reversed_tensors = (
+            self._reverse_check_min_max_values or
+            self._reverse_check_finite or
+            self._reverse_keep_tensors
+        )
+        ret = self._reverse_model(
+            model,
+            stop_analysis_at_tensors=stop_analysis_at_tensors,
+            return_all_reversed_tensors=return_all_reversed_tensors)
+
+        if return_all_reversed_tensors:
+            ret = (self._postprocess_analysis(ret[0]), ret[1])
+        else:
+            ret = self._postprocess_analysis(ret)
+
+        if return_all_reversed_tensors:
+            debug_tensors = []
+            self._debug_tensors_indices = {}
+
+            values = list(six.itervalues(ret[1]))
+            mapping = {i: v["id"] for i, v in enumerate(values)}
+            tensors = [v["final_tensor"] for v in values]
+            self._reverse_tensors_mapping = mapping
+
+            if self._reverse_check_min_max_values:
+                tmp = [ilayers.Min(None)(x) for x in tensors]
+                self._debug_tensors_indices["min"] = (
+                    len(debug_tensors),
+                    len(debug_tensors)+len(tmp))
+                debug_tensors += tmp
+
+                tmp = [ilayers.Max(None)(x) for x in tensors]
+                self._debug_tensors_indices["max"] = (
+                    len(debug_tensors),
+                    len(debug_tensors)+len(tmp))
+                debug_tensors += tmp
+
+            if self._reverse_check_finite:
+                tmp = iutils.to_list(ilayers.FiniteCheck()(tensors))
+                self._debug_tensors_indices["finite"] = (
+                    len(debug_tensors),
+                    len(debug_tensors)+len(tmp))
+                debug_tensors += tmp
+
+            if self._reverse_keep_tensors:
+                self._debug_tensors_indices["keep"] = (
+                    len(debug_tensors),
+                    len(debug_tensors)+len(tensors))
+                debug_tensors += tensors
+
+            ret = (ret[0], debug_tensors)
+        return ret
+
+    def _handle_debug_output(self, debug_values):
+
+        if self._reverse_check_min_max_values:
+            indices = self._debug_tensors_indices["min"]
+            tmp = debug_values[indices[0]:indices[1]]
+            tmp = sorted([(self._reverse_tensors_mapping[i], v)
+                          for i, v in enumerate(tmp)])
+            print("Minimum values in tensors: "
+                  "((NodeID, TensorID), Value) - {}".format(tmp))
+
+            indices = self._debug_tensors_indices["max"]
+            tmp = debug_values[indices[0]:indices[1]]
+            tmp = sorted([(self._reverse_tensors_mapping[i], v)
+                          for i, v in enumerate(tmp)])
+            print("Maximum values in tensors: "
+                  "((NodeID, TensorID), Value) - {}".format(tmp))
+
+        if self._reverse_check_finite:
+            indices = self._debug_tensors_indices["finite"]
+            tmp = debug_values[indices[0]:indices[1]]
+            nfinite_tensors = np.flatnonzero(np.asarray(tmp) > 0)
+
+            if len(nfinite_tensors) > 0:
+                nfinite_tensors = sorted([self._reverse_tensors_mapping[i]
+                                          for i in nfinite_tensors])
+                print("Not finite values found in following nodes: "
+                      "(NodeID, TensorID) - {}".format(nfinite_tensors))
+
+        if self._reverse_keep_tensors:
+            indices = self._debug_tensors_indices["keep"]
+            tmp = debug_values[indices[0]:indices[1]]
+            tmp = sorted([(self._reverse_tensors_mapping[i], v)
+                          for i, v in enumerate(tmp)])
+            self._reversed_tensors = tmp
+
+    def _get_state(self):
+        state = super(ReverseAnalyzerBase, self)._get_state()
+        state.update({"reverse_verbose": self._reverse_verbose})
+        state.update({"reverse_clip_values": self._reverse_clip_values})
+        state.update({"reverse_project_bottleneck_layers":
+                      self._reverse_project_bottleneck_layers})
+        state.update({"reverse_check_min_max_values":
+                      self._reverse_check_min_max_values})
+        state.update({"reverse_check_finite": self._reverse_check_finite})
+        state.update({"reverse_keep_tensors": self._reverse_keep_tensors})
+        state.update({"reverse_reapply_on_copied_layers":
+                      self._reverse_reapply_on_copied_layers})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        reverse_verbose = state.pop("reverse_verbose")
+        reverse_clip_values = state.pop("reverse_clip_values")
+        reverse_project_bottleneck_layers = (
+            state.pop("reverse_project_bottleneck_layers"))
+        reverse_check_min_max_values = (
+            state.pop("reverse_check_min_max_values"))
+        reverse_check_finite = state.pop("reverse_check_finite")
+        reverse_keep_tensors = state.pop("reverse_keep_tensors")
+        reverse_reapply_on_copied_layers = (
+            state.pop("reverse_reapply_on_copied_layers"))
+        kwargs = super(ReverseAnalyzerBase, clazz)._state_to_kwargs(state)
+        kwargs.update({"reverse_verbose": reverse_verbose,
+                       "reverse_clip_values": reverse_clip_values,
+                       "reverse_project_bottleneck_layers":
+                       reverse_project_bottleneck_layers,
+                       "reverse_check_min_max_values":
+                       reverse_check_min_max_values,
+                       "reverse_check_finite": reverse_check_finite,
+                       "reverse_keep_tensors": reverse_keep_tensors,
+                       "reverse_reapply_on_copied_layers":
+                       reverse_reapply_on_copied_layers})
+        return kwargs

+ 405 - 0
original_model/innvestigate/analyzer/deeplift.py

@@ -0,0 +1,405 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import importlib
+import keras.backend as K
+import keras.layers
+import numpy as np
+import tempfile
+import warnings
+
+
+from . import base
+from .. import layers as ilayers
+from .. import utils as iutils
+from ..utils import keras as kutils
+from ..utils.keras import checks as kchecks
+from ..utils.keras import graph as kgraph
+
+
+__all__ = [
+    "DeepLIFT",
+    "DeepLIFTWrapper",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def _create_deeplift_rules(reference_mapping, approximate_gradient=True):
+    def RescaleRule(Xs, Ys, As, reverse_state, local_references={}):
+        if approximate_gradient:
+            def rescale_f(x):
+                a, dx, dy, g = x
+                return K.switch(K.less(K.abs(dx), K.epsilon()), g, a*(dy/dx))
+        else:
+            def rescale_f(x):
+                a, dx, dy, _ = x
+                return a*(dy/(dx + K.epsilon()))
+
+        grad = ilayers.GradientWRT(len(Xs))
+        rescale = keras.layers.Lambda(rescale_f)
+
+        Xs_references = [
+            reference_mapping.get(x, local_references.get(x, None))
+            for x in Xs
+        ]
+        Ys_references = [
+            reference_mapping.get(x, local_references.get(x, None))
+            for x in Ys
+        ]
+
+        Xs_differences = [keras.layers.Subtract()([x, r])
+                          for x, r in zip(Xs, Xs_references)]
+        Ys_differences = [keras.layers.Subtract()([x, r])
+                          for x, r in zip(Ys, Ys_references)]
+        gradients = iutils.to_list(grad(Xs+Ys+As))
+
+        return [rescale([a, dx, dy, g])
+                for a, dx, dy, g
+                in zip(As, Xs_differences, Ys_differences, gradients)]
+
+    def LinearRule(Xs, Ys, As, reverse_state):
+        if approximate_gradient:
+            def switch_f(x):
+                dx, a, g = x
+                return K.switch(K.less(K.abs(dx), K.epsilon()), g, a)
+        else:
+            def switch_f(x):
+                _, a, _ = x
+                return a
+
+        grad = ilayers.GradientWRT(len(Xs))
+        switch = keras.layers.Lambda(switch_f)
+
+        Xs_references = [reference_mapping[x] for x in Xs]
+
+        Ys_references = [reference_mapping[x] for x in Ys]
+
+        Xs_differences = [keras.layers.Subtract()([x, r])
+                          for x, r in zip(Xs, Xs_references)]
+        Ys_differences = [keras.layers.Subtract()([x, r])
+                          for x, r in zip(Ys, Ys_references)]
+
+        # Divide incoming relevance by the activations.
+        tmp = [ilayers.SafeDivide()([a, b])
+               for a, b in zip(As, Ys_differences)]
+
+        # Propagate the relevance to input neurons
+        # using the gradient.
+        tmp = iutils.to_list(grad(Xs+Ys+tmp))
+
+        # Re-weight relevance with the input values.
+        tmp = [keras.layers.Multiply()([a, b])
+               for a, b in zip(Xs_differences, tmp)]
+
+        # only the gradient
+        gradients = iutils.to_list(grad(Xs+Ys+As))
+
+        return [switch([dx, a, g])
+                for dx, a, g
+                in zip(Xs_differences, tmp, gradients)]
+
+    return RescaleRule, LinearRule
+
+
+class DeepLIFT(base.ReverseAnalyzerBase):
+    """DeepLIFT-rescale algorithm
+
+    This class implements the DeepLIFT algorithm using
+    the rescale rule (as in DeepExplain (Ancona et.al.)).
+
+    WARNING: This implementation contains bugs.
+
+    :param model: A Keras model.
+    """
+
+    def __init__(self, model, *args, **kwargs):
+        warnings.warn("This implementation contains bugs.")
+        self._reference_inputs = kwargs.pop("reference_inputs", 0)
+        self._approximate_gradient = kwargs.pop(
+            "approximate_gradient", True)
+        self._add_model_softmax_check()
+        super(DeepLIFT, self).__init__(model, *args, **kwargs)
+
+    def _prepare_model(self, model):
+        ret = super(DeepLIFT, self)._prepare_model(model)
+        # Store analysis input to create reference inputs.
+        self._analysis_inputs = ret[1]
+        return ret
+
+    def _create_reference_activations(self, model):
+        self._model_execution_trace = kgraph.trace_model_execution(model)
+        layers, execution_list, outputs = self._model_execution_trace
+
+        self._reference_activations = {}
+
+        # Create references and graph inputs.
+        tmp = kutils.broadcast_np_tensors_to_keras_tensors(
+            model.inputs, self._reference_inputs)
+        tmp = [K.variable(x) for x in tmp]
+
+        constant_reference_inputs = [
+            keras.layers.Input(tensor=x, shape=K.int_shape(x)[1:])
+            for x in tmp
+        ]
+
+        for k, v in zip(model.inputs, constant_reference_inputs):
+            self._reference_activations[k] = v
+
+        for k, v in zip(self._analysis_inputs, self._analysis_inputs):
+            self._reference_activations[k] = v
+
+        # Compute intermediate states.
+        for layer, Xs, Ys in execution_list:
+            activations = [self._reference_activations[x] for x in Xs]
+
+            if isinstance(layer, keras.layers.InputLayer):
+                # Special case. Do nothing.
+                next_activations = activations
+            else:
+                next_activations = iutils.to_list(
+                    kutils.apply(layer, activations))
+
+            assert len(next_activations) == len(Ys)
+            for k, v in zip(Ys, next_activations):
+                self._reference_activations[k] = v
+
+        return constant_reference_inputs
+
+    def _create_analysis(self, model, *args, **kwargs):
+        constant_reference_inputs = self._create_reference_activations(model)
+
+        RescaleRule, LinearRule = _create_deeplift_rules(
+            self._reference_activations, self._approximate_gradient)
+
+        # Kernel layers.
+        self._add_conditional_reverse_mapping(
+            lambda l: kchecks.contains_kernel(l),
+            LinearRule,
+            name="deeplift_kernel_layer",
+        )
+
+        # Activation layers
+        self._add_conditional_reverse_mapping(
+            lambda l: (not kchecks.contains_kernel(l) and
+                       kchecks.contains_activation(l)),
+            RescaleRule,
+            name="deeplift_activation_layer",
+        )
+
+        tmp = super(DeepLIFT, self)._create_analysis(
+            model, *args, **kwargs)
+
+        if isinstance(tmp, tuple):
+            if len(tmp) == 3:
+                analysis_outputs, debug_outputs, constant_inputs = tmp
+            elif len(tmp) == 2:
+                analysis_outputs, debug_outputs = tmp
+                constant_inputs = list()
+            elif len(tmp) == 1:
+                analysis_outputs = iutils.to_list(tmp[0])
+                constant_inputs, debug_outputs = list(), list()
+            else:
+                raise Exception("Unexpected output from _create_analysis.")
+        else:
+            analysis_outputs = tmp
+            constant_inputs, debug_outputs = list(), list()
+
+        return (analysis_outputs,
+                debug_outputs,
+                constant_inputs+constant_reference_inputs)
+
+    def _head_mapping(self, X):
+        return keras.layers.Subtract()([X, self._reference_activations[X]])
+
+    def _reverse_model(self,
+                       model,
+                       stop_analysis_at_tensors=[],
+                       return_all_reversed_tensors=False):
+        return kgraph.reverse_model(
+            model,
+            reverse_mappings=self._reverse_mapping,
+            default_reverse_mapping=self._default_reverse_mapping,
+            head_mapping=self._head_mapping,
+            stop_mapping_at_tensors=stop_analysis_at_tensors,
+            verbose=self._reverse_verbose,
+            clip_all_reversed_tensors=self._reverse_clip_values,
+            project_bottleneck_tensors=self._reverse_project_bottleneck_layers,
+            return_all_reversed_tensors=return_all_reversed_tensors,
+            execution_trace=self._model_execution_trace)
+
+    def _get_state(self):
+        state = super(DeepLIFT, self)._get_state()
+        state.update({"reference_inputs": self._reference_inputs})
+        state.update({"approximate_gradient": self._approximate_gradient})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        reference_inputs = state.pop("reference_inputs")
+        approximate_gradient = state.pop("approximate_gradient")
+        kwargs = super(DeepLIFT, clazz)._state_to_kwargs(state)
+        kwargs.update({"reference_inputs": reference_inputs})
+        kwargs.update({"approximate_gradient": approximate_gradient})
+        return kwargs
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class DeepLIFTWrapper(base.AnalyzerNetworkBase):
+    """Wrapper around DeepLIFT package
+
+    This class wraps the DeepLIFT package.
+    For further explanation of the parameters check out:
+    https://github.com/kundajelab/deeplift
+
+    :param model: A Keras model.
+    :param nonlinear_mode: The nonlinear mode parameter.
+    :param reference_inputs: The reference input used for DeepLIFT.
+    :param verbose: Verbosity of the DeepLIFT package.
+
+    :note: Requires the deeplift package.
+    """
+
+    def __init__(self, model, **kwargs):
+        self._nonlinear_mode = kwargs.pop("nonlinear_mode", "rescale")
+        self._reference_inputs = kwargs.pop("reference_inputs", 0)
+        self._verbose = kwargs.pop("verbose", False)
+        #relevant for "index" selection mode
+        self._batch_size = kwargs.pop("batch_size", 32)
+        self._add_model_softmax_check()
+
+        try:
+            self._deeplift_module = importlib.import_module("deeplift")
+        except ImportError:
+            raise ImportError("To use DeepLIFTWrapper please install "
+                              "the python module 'deeplift', e.g.: "
+                              "'pip install deeplift'")
+
+        super(DeepLIFTWrapper, self).__init__(model, **kwargs)
+
+    def _create_deep_lift_func(self):
+        # Store model and load into deeplift format.
+        kc = importlib.import_module("deeplift.conversion.kerasapi_conversion")
+        modes = self._deeplift_module.layers.NonlinearMxtsMode
+
+        key = self._nonlinear_mode
+        nonlinear_mxts_mode = {
+            "genomics_default": modes.DeepLIFT_GenomicsDefault,
+            "reveal_cancel": modes.RevealCancel,
+            "rescale": modes.Rescale,
+        }[key]
+
+        with tempfile.NamedTemporaryFile(suffix=".hdf5") as f:
+            self._model.save(f.name)
+            deeplift_model = kc.convert_model_from_saved_files(
+                f.name, nonlinear_mxts_mode=nonlinear_mxts_mode,
+                verbose=self._verbose)
+
+        # Create function with respect to input layers
+        def fix_name(s):
+            return s.replace(":", "_")
+
+        score_layer_names = [fix_name(l.name) for l in self._model.inputs]
+        if len(self._model.outputs) > 1:
+            raise ValueError("Only a single output layer is supported.")
+        tmp = self._model.outputs[0]._keras_history
+        target_layer_name = fix_name(tmp[0].name+"_%i" % tmp[1])
+        self._func = deeplift_model.get_target_contribs_func(
+            find_scores_layer_name=score_layer_names,
+            pre_activation_target_layer_name=target_layer_name)
+        self._references = kutils.broadcast_np_tensors_to_keras_tensors(
+            self._model.inputs, self._reference_inputs)
+
+    def _analyze_with_deeplift(self, X, neuron_idx, batch_size):
+        return self._func(task_idx=neuron_idx,
+                          input_data_list=X,
+                          batch_size=batch_size,
+                          input_references_list=self._references,
+                          progress_update=None)
+
+    def analyze(self, X, neuron_selection=None):
+        if not hasattr(self, "_func"):
+            self._create_deep_lift_func()
+
+        X = iutils.to_list(X)
+
+        if(neuron_selection is not None and
+           self._neuron_selection_mode != "index"):
+            raise ValueError("Only neuron_selection_mode 'index' expects "
+                             "the neuron_selection parameter.")
+        if(neuron_selection is None and
+           self._neuron_selection_mode == "index"):
+            raise ValueError("neuron_selection_mode 'index' expects "
+                             "the neuron_selection parameter.")
+
+        if self._neuron_selection_mode == "index":
+            neuron_selection = np.asarray(neuron_selection).flatten()
+            if neuron_selection.size != 1:
+                # The code allows to select multiple neurons.
+                raise ValueError("One neuron can be selected with DeepLIFT.")
+
+            neuron_idx = neuron_selection[0]
+            analysis = self._analyze_with_deeplift(X, neuron_idx, self._batch_size)
+
+            # Parse the output.
+            ret = []
+            for x, analysis_for_x in zip(X, analysis):
+                tmp = np.stack([a for a in analysis_for_x])
+                tmp = tmp.reshape(x.shape)
+                ret.append(tmp)
+        elif self._neuron_selection_mode == "max_activation":
+            neuron_idx = np.argmax(self._model.predict_on_batch(X), axis=1)
+
+            analysis = []
+            # run for each batch with its respective max activated neuron
+            for i, ni in enumerate(neuron_idx):
+                # slice input tensors
+                tmp = [x[i:i+1] for x in X]
+                analysis.append(self._analyze_with_deeplift(tmp, ni, 1))
+
+            # Parse the output.
+            ret = []
+            for i, x in enumerate(X):
+                tmp = np.stack([a[i] for a in analysis]).reshape(x.shape)
+                ret.append(tmp)
+        else:
+            raise ValueError("Only neuron_selection_mode index or "
+                             "max_activation are supported.")
+
+        if isinstance(ret, list) and len(ret) == 1:
+            ret = ret[0]
+        return ret
+
+    def _get_state(self):
+        state = super(DeepLIFTWrapper, self)._get_state()
+        state.update({"nonlinear_mode": self._nonlinear_mode})
+        state.update({"reference_inputs": self._reference_inputs})
+        state.update({"verbose": self._verbose})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        nonlinear_mode = state.pop("nonlinear_mode")
+        reference_inputs = state.pop("reference_inputs")
+        verbose = state.pop("verbose")
+        kwargs = super(DeepLIFTWrapper, clazz)._state_to_kwargs(state)
+        kwargs.update({
+            "nonlinear_mode": nonlinear_mode,
+            "reference_inputs": reference_inputs,
+            "verbose": verbose,
+        })
+        return kwargs

+ 204 - 0
original_model/innvestigate/analyzer/deeptaylor.py

@@ -0,0 +1,204 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.layers
+import keras.models
+
+
+from . import base
+from .relevance_based import relevance_rule as lrp_rules
+from ..utils.keras import checks as kchecks
+from ..utils.keras import graph as kgraph
+
+
+__all__ = [
+    "DeepTaylor",
+    "BoundedDeepTaylor",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class DeepTaylor(base.ReverseAnalyzerBase):
+    """DeepTaylor for ReLU-networks with unbounded input
+
+    This class implements the DeepTaylor algorithm for neural networks with
+    ReLU activation and unbounded input ranges.
+
+    :param model: A Keras model.
+    """
+
+    def __init__(self, model, *args, **kwargs):
+
+        self._add_model_softmax_check()
+        self._add_model_check(
+            lambda layer: not kchecks.only_relu_activation(layer),
+            "This DeepTaylor implementation only supports ReLU activations.",
+            check_type="exception",
+        )
+        super(DeepTaylor, self).__init__(model, *args, **kwargs)
+
+    def _create_analysis(self, *args, **kwargs):
+
+        def do_nothing(Xs, Ys, As, reverse_state):
+            return As
+
+        # Kernel layers.
+        self._add_conditional_reverse_mapping(
+            lambda l: (kchecks.contains_kernel(l) and
+                       kchecks.contains_activation(l)),
+            lrp_rules.Alpha1Beta0IgnoreBiasRule,
+            name="deep_taylor_kernel_w_relu",
+        )
+        self._add_conditional_reverse_mapping(
+            lambda l: (kchecks.contains_kernel(l) and
+                       not kchecks.contains_activation(l)),
+            lrp_rules.WSquareRule,
+            name="deep_taylor_kernel_wo_relu",
+        )
+
+        # ReLU Activation layer
+        self._add_conditional_reverse_mapping(
+            lambda l: (not kchecks.contains_kernel(l) and
+                       kchecks.contains_activation(l)),
+            self._gradient_reverse_mapping,
+            name="deep_taylor_relu",
+        )
+
+        # Assume conv layer beforehand -> unbounded
+        bn_mapping = kgraph.apply_mapping_to_fused_bn_layer(
+            lrp_rules.WSquareRule,
+            fuse_mode="one_linear",
+        )
+        self._add_conditional_reverse_mapping(
+            kchecks.is_batch_normalization_layer,
+            bn_mapping,
+            name="deep_taylor_batch_norm",
+        )
+        # Special layers.
+        self._add_conditional_reverse_mapping(
+            kchecks.is_max_pooling,
+            self._gradient_reverse_mapping,
+            name="deep_taylor_max_pooling",
+        )
+        self._add_conditional_reverse_mapping(
+            kchecks.is_average_pooling,
+            self._gradient_reverse_mapping,
+            name="deep_taylor_average_pooling",
+        )
+        self._add_conditional_reverse_mapping(
+            lambda l: isinstance(l, keras.layers.Add),
+            # Ignore scaling with 0.5
+            self._gradient_reverse_mapping,
+            name="deep_taylor_add",
+        )
+        self._add_conditional_reverse_mapping(
+            lambda l: isinstance(l, (
+                keras.layers.convolutional.UpSampling1D,
+                keras.layers.convolutional.UpSampling2D,
+                keras.layers.convolutional.UpSampling3D,
+                keras.layers.core.Dropout,
+                keras.layers.core.SpatialDropout1D,
+                keras.layers.core.SpatialDropout2D,
+                keras.layers.core.SpatialDropout3D,
+            )),
+            self._gradient_reverse_mapping,
+            name="deep_taylor_special_layers",
+        )
+
+        # Layers w/o transformation
+        self._add_conditional_reverse_mapping(
+            lambda l: isinstance(l, (
+                keras.engine.topology.InputLayer,
+                keras.layers.convolutional.Cropping1D,
+                keras.layers.convolutional.Cropping2D,
+                keras.layers.convolutional.Cropping3D,
+                keras.layers.convolutional.ZeroPadding1D,
+                keras.layers.convolutional.ZeroPadding2D,
+                keras.layers.convolutional.ZeroPadding3D,
+                keras.layers.Concatenate,
+                keras.layers.core.Flatten,
+                keras.layers.core.Masking,
+                keras.layers.core.Permute,
+                keras.layers.core.RepeatVector,
+                keras.layers.core.Reshape,
+            )),
+            self._gradient_reverse_mapping,
+            name="deep_taylor_no_transform",
+        )
+
+        return super(DeepTaylor, self)._create_analysis(
+            *args, **kwargs)
+
+    def _default_reverse_mapping(self, Xs, Ys, reversed_Ys, reverse_state):
+        """
+        Block all default mappings.
+        """
+        raise NotImplementedError(
+            "Layer %s not supported." % reverse_state["layer"])
+
+    def _prepare_model(self, model):
+        """
+        To be theoretically sound Deep-Taylor expects only positive outputs.
+        """
+
+        positive_outputs = [keras.layers.ReLU()(x) for x in model.outputs]
+        model_with_positive_output = keras.models.Model(
+            inputs=model.inputs, outputs=positive_outputs)
+
+        return super(DeepTaylor, self)._prepare_model(
+            model_with_positive_output)
+
+
+class BoundedDeepTaylor(DeepTaylor):
+    """DeepTaylor for ReLU-networks with bounded input
+
+    This class implements the DeepTaylor algorithm for neural networks with
+    ReLU activation and bounded input ranges.
+
+    :param model: A Keras model.
+    :param low: Lowest value of the input range. See Z_B rule.
+    :param high: Highest value of the input range. See Z_B rule.
+    """
+
+    def __init__(self, model, low=None, high=None, **kwargs):
+
+        if low is None or high is None:
+            raise ValueError("The low or high parameter is missing."
+                             " Z-B (bounded rule) require both values.")
+
+        self._bounds_low = low
+        self._bounds_high = high
+
+        super(BoundedDeepTaylor, self).__init__(
+            model, **kwargs)
+
+    def _create_analysis(self, *args, **kwargs):
+
+        low, high = self._bounds_low, self._bounds_high
+
+        class BoundedProxyRule(lrp_rules.BoundedRule):
+            def __init__(self, *args, **kwargs):
+                super(BoundedProxyRule, self).__init__(
+                    *args, low=low, high=high,
+                    **kwargs)
+
+        self._add_conditional_reverse_mapping(
+            lambda l: kchecks.is_input_layer(l) and kchecks.contains_kernel(l),
+            BoundedProxyRule,
+            name="deep_taylor_first_layer_bounded",
+            priority=10,  # do first
+        )
+
+        return super(BoundedDeepTaylor, self)._create_analysis(
+            *args, **kwargs)

+ 310 - 0
original_model/innvestigate/analyzer/gradient_based.py

@@ -0,0 +1,310 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.models
+import keras
+
+
+from . import base
+from . import wrapper
+from .. import layers as ilayers
+from .. import utils as iutils
+from ..utils import keras as kutils
+from ..utils.keras import checks as kchecks
+from ..utils.keras import graph as kgraph
+
+__all__ = [
+    "BaselineGradient",
+    "Gradient",
+
+    "InputTimesGradient",
+
+    "Deconvnet",
+    "GuidedBackprop",
+
+    "IntegratedGradients",
+
+    "SmoothGrad",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class BaselineGradient(base.AnalyzerNetworkBase):
+    """Gradient analyzer based on build-in gradient.
+
+    Returns as analysis the function value with respect to the input.
+    The gradient is computed via the build in function.
+    Is mainly used for debugging purposes.
+
+    :param model: A Keras model.
+    """
+
+    def __init__(self, model, postprocess=None, **kwargs):
+
+        if postprocess not in [None, "abs", "square"]:
+            raise ValueError("Parameter 'postprocess' must be either "
+                             "None, 'abs', or 'square'.")
+        self._postprocess = postprocess
+
+        self._add_model_softmax_check()
+
+        super(BaselineGradient, self).__init__(model, **kwargs)
+
+    def _create_analysis(self, model, stop_analysis_at_tensors=[]):
+        tensors_to_analyze = [x for x in iutils.to_list(model.inputs)
+                              if x not in stop_analysis_at_tensors]
+        ret = iutils.to_list(ilayers.Gradient()(
+            tensors_to_analyze+[model.outputs[0]]))
+
+        if self._postprocess == "abs":
+            ret = ilayers.Abs()(ret)
+        elif self._postprocess == "square":
+            ret = ilayers.Square()(ret)
+
+        return iutils.to_list(ret)
+
+    def _get_state(self):
+        state = super(BaselineGradient, self)._get_state()
+        state.update({"postprocess": self._postprocess})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        postprocess = state.pop("postprocess")
+        kwargs = super(BaselineGradient, clazz)._state_to_kwargs(state)
+        kwargs.update({
+            "postprocess": postprocess,
+        })
+        return kwargs
+
+
+class Gradient(base.ReverseAnalyzerBase):
+    """Gradient analyzer.
+
+    Returns as analysis the function value with respect to the input.
+    The gradient is computed via the librarie's network reverting.
+
+    :param model: A Keras model.
+    """
+
+    def __init__(self, model, postprocess=None, **kwargs):
+
+        if postprocess not in [None, "abs", "square"]:
+            raise ValueError("Parameter 'postprocess' must be either "
+                             "None, 'abs', or 'square'.")
+        self._postprocess = postprocess
+
+        self._add_model_softmax_check()
+
+        super(Gradient, self).__init__(model, **kwargs)
+
+    def _head_mapping(self, X):
+        return ilayers.OnesLike()(X)
+
+    def _postprocess_analysis(self, X):
+        ret = super(Gradient, self)._postprocess_analysis(X)
+
+        if self._postprocess == "abs":
+            ret = ilayers.Abs()(ret)
+        elif self._postprocess == "square":
+            ret = ilayers.Square()(ret)
+
+        return iutils.to_list(ret)
+
+    def _get_state(self):
+        state = super(Gradient, self)._get_state()
+        state.update({"postprocess": self._postprocess})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        postprocess = state.pop("postprocess")
+        kwargs = super(Gradient, clazz)._state_to_kwargs(state)
+        kwargs.update({
+            "postprocess": postprocess,
+        })
+        return kwargs
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class InputTimesGradient(Gradient):
+    """Input*Gradient analyzer.
+
+    :param model: A Keras model.
+    """
+
+    def __init__(self, model, **kwargs):
+
+        self._add_model_softmax_check()
+
+        super(InputTimesGradient, self).__init__(model, **kwargs)
+
+    def _create_analysis(self, model, stop_analysis_at_tensors=[]):
+        tensors_to_analyze = [x for x in iutils.to_list(model.inputs)
+                              if x not in stop_analysis_at_tensors]
+        gradients = super(InputTimesGradient, self)._create_analysis(
+            model, stop_analysis_at_tensors=stop_analysis_at_tensors)
+        return [keras.layers.Multiply()([i, g])
+                for i, g in zip(tensors_to_analyze, gradients)]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class DeconvnetReverseReLULayer(kgraph.ReverseMappingBase):
+
+    def __init__(self, layer, state):
+        self._activation = keras.layers.Activation("relu")
+        self._layer_wo_relu = kgraph.copy_layer_wo_activation(
+            layer,
+            name_template="reversed_%s",
+        )
+
+    def apply(self, Xs, Ys, reversed_Ys, reverse_state):
+        # Apply relus conditioned on backpropagated values.
+        reversed_Ys = kutils.apply(self._activation, reversed_Ys)
+
+        # Apply gradient of forward pass without relus.
+        Ys_wo_relu = kutils.apply(self._layer_wo_relu, Xs)
+        return ilayers.GradientWRT(len(Xs))(Xs+Ys_wo_relu+reversed_Ys)
+
+
+class Deconvnet(base.ReverseAnalyzerBase):
+    """Deconvnet analyzer.
+
+    Applies the "deconvnet" algorithm to analyze the model.
+
+    :param model: A Keras model.
+    """
+
+    def __init__(self, model, **kwargs):
+
+        self._add_model_softmax_check()
+        self._add_model_check(
+            lambda layer: not kchecks.only_relu_activation(layer),
+            "Deconvnet is only specified for networks with ReLU activations.",
+            check_type="exception",
+        )
+
+        super(Deconvnet, self).__init__(model, **kwargs)
+
+    def _create_analysis(self, *args, **kwargs):
+
+        self._add_conditional_reverse_mapping(
+            lambda layer: kchecks.contains_activation(layer, "relu"),
+            DeconvnetReverseReLULayer,
+            name="deconvnet_reverse_relu_layer",
+        )
+
+        return super(Deconvnet, self)._create_analysis(*args, **kwargs)
+
+
+def GuidedBackpropReverseReLULayer(Xs, Ys, reversed_Ys, reverse_state):
+    activation = keras.layers.Activation("relu")
+    # Apply relus conditioned on backpropagated values.
+    reversed_Ys = kutils.apply(activation, reversed_Ys)
+
+    # Apply gradient of forward pass.
+    return ilayers.GradientWRT(len(Xs))(Xs+Ys+reversed_Ys)
+
+
+class GuidedBackprop(base.ReverseAnalyzerBase):
+    """Guided backprop analyzer.
+
+    Applies the "guided backprop" algorithm to analyze the model.
+
+    :param model: A Keras model.
+    """
+
+    def __init__(self, model, **kwargs):
+
+        self._add_model_softmax_check()
+        self._add_model_check(
+            lambda layer: not kchecks.only_relu_activation(layer),
+            "GuidedBackprop is only specified for "
+            "networks with ReLU activations.",
+            check_type="exception",
+        )
+
+        super(GuidedBackprop, self).__init__(model, **kwargs)
+
+    def _create_analysis(self, *args, **kwargs):
+
+        self._add_conditional_reverse_mapping(
+            lambda layer: kchecks.contains_activation(layer, "relu"),
+            GuidedBackpropReverseReLULayer,
+            name="guided_backprop_reverse_relu_layer",
+        )
+
+        return super(GuidedBackprop, self)._create_analysis(*args, **kwargs)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class IntegratedGradients(wrapper.PathIntegrator):
+    """Integrated gradient analyzer.
+
+    Applies the "integrated gradient" algorithm to analyze the model.
+
+    :param model: A Keras model.
+    :param steps: Number of steps to use average along integration path.
+    """
+
+    def __init__(self, model, steps=64, **kwargs):
+        subanalyzer_kwargs = {}
+        kwargs_keys = ["neuron_selection_mode", "postprocess"]
+        for key in kwargs_keys:
+            if key in kwargs:
+                subanalyzer_kwargs[key] = kwargs.pop(key)
+        subanalyzer = Gradient(model, **subanalyzer_kwargs)
+
+        super(IntegratedGradients, self).__init__(subanalyzer,
+                                                  steps=steps,
+                                                  **kwargs)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class SmoothGrad(wrapper.GaussianSmoother):
+    """Smooth grad analyzer.
+
+    Applies the "smooth grad" algorithm to analyze the model.
+
+    :param model: A Keras model.
+    :param augment_by_n: Number of distortions to average for smoothing.
+    """
+
+    def __init__(self, model, augment_by_n=64, **kwargs):
+        subanalyzer_kwargs = {}
+        kwargs_keys = ["neuron_selection_mode", "postprocess"]
+        for key in kwargs_keys:
+            if key in kwargs:
+                subanalyzer_kwargs[key] = kwargs.pop(key)
+        subanalyzer = Gradient(model, **subanalyzer_kwargs)
+
+        super(SmoothGrad, self).__init__(subanalyzer,
+                                         augment_by_n=augment_by_n,
+                                         **kwargs)

+ 68 - 0
original_model/innvestigate/analyzer/misc.py

@@ -0,0 +1,68 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+from .base import AnalyzerNetworkBase
+from .. import layers as ilayers
+from .. import utils as iutils
+
+
+__all__ = ["Random", "Input"]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class Input(AnalyzerNetworkBase):
+    """Returns the input.
+
+    Returns the input as analysis.
+
+    :param model: A Keras model.
+    """
+
+    def _create_analysis(self, model, stop_analysis_at_tensors=[]):
+        tensors_to_analyze = [x for x in iutils.to_list(model.inputs)
+                              if x not in stop_analysis_at_tensors]
+        return [ilayers.Identity()(x) for x in tensors_to_analyze]
+
+
+class Random(AnalyzerNetworkBase):
+    """Returns noise.
+
+    Returns the Gaussian noise as analysis.
+
+    :param model: A Keras model.
+    :param stddev: The standard deviation of the noise.
+    """
+
+    def __init__(self, model, stddev=1, **kwargs):
+        self._stddev = stddev
+
+        super(Random, self).__init__(model, **kwargs)
+
+    def _create_analysis(self, model, stop_analysis_at_tensors=[]):
+        noise = ilayers.TestPhaseGaussianNoise(stddev=self._stddev)
+        tensors_to_analyze = [x for x in iutils.to_list(model.inputs)
+                              if x not in stop_analysis_at_tensors]
+        return [noise(x) for x in tensors_to_analyze]
+
+    def _get_state(self):
+        state = super(Random, self)._get_state()
+        state.update({"stddev": self._stddev})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        stddev = state.pop("stddev")
+        kwargs = super(Random, clazz)._state_to_kwargs(state)
+        kwargs.update({"stddev": stddev})
+        return kwargs

+ 272 - 0
original_model/innvestigate/analyzer/pattern_based.py

@@ -0,0 +1,272 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+import keras.activations
+import keras.engine.topology
+import keras.layers
+import keras.layers.core
+import keras.layers.pooling
+import keras.models
+import keras
+import numpy as np
+import warnings
+
+
+from . import base
+from .. import layers as ilayers
+from .. import utils
+from .. import tools as itools
+from ..utils import keras as kutils
+from ..utils.keras import checks as kchecks
+from ..utils.keras import graph as kgraph
+
+
+__all__ = [
+    "PatternNet",
+    "PatternAttribution",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+SUPPORTED_LAYER_PATTERNNET = (
+    keras.engine.topology.InputLayer,
+    keras.layers.convolutional.Conv2D,
+    keras.layers.core.Dense,
+    keras.layers.core.Dropout,
+    keras.layers.core.Flatten,
+    keras.layers.core.Masking,
+    keras.layers.core.Permute,
+    keras.layers.core.Reshape,
+    keras.layers.Concatenate,
+    keras.layers.pooling.GlobalMaxPooling1D,
+    keras.layers.pooling.GlobalMaxPooling2D,
+    keras.layers.pooling.GlobalMaxPooling3D,
+    keras.layers.pooling.MaxPooling1D,
+    keras.layers.pooling.MaxPooling2D,
+    keras.layers.pooling.MaxPooling3D,
+)
+
+
+class PatternNetReverseKernelLayer(kgraph.ReverseMappingBase):
+    """
+    PatternNet backward mapping for layers with kernels.
+
+    Applies the (filter) weights on the forward pass and
+    on the backward pass applies the gradient computation
+    where the filter weights are replaced with the patterns.
+    """
+
+    def __init__(self, layer, state, pattern):
+        config = layer.get_config()
+
+        # Layer can contain a kernel and an activation.
+        # Split layers in a kernel layer and an activation layer.
+        activation = None
+        if "activation" in config:
+            activation = config["activation"]
+            config["activation"] = None
+        self._act_layer = keras.layers.Activation(
+            activation,
+            name="reversed_act_%s" % config["name"])
+        self._filter_layer = kgraph.copy_layer_wo_activation(
+            layer, name_template="reversed_filter_%s")
+
+        # Replace filter/kernel weights with patterns.
+        filter_weights = layer.get_weights()
+        # Assume that only one weight has a corresponding pattern.
+        # E.g., biases have no pattern.
+        tmp = [pattern.shape == x.shape for x in filter_weights]
+        if np.sum(tmp) != 1:
+            raise Exception("Cannot match pattern to filter.")
+        filter_weights[np.argmax(tmp)] = pattern
+        self._pattern_layer = kgraph.copy_layer_wo_activation(
+            layer,
+            name_template="reversed_pattern_%s",
+            weights=filter_weights)
+
+    def apply(self, Xs, Ys, reversed_Ys, reverse_state):
+        # Reapply the prepared layers.
+        act_Xs = kutils.apply(self._filter_layer, Xs)
+        act_Ys = kutils.apply(self._act_layer, act_Xs)
+        pattern_Ys = kutils.apply(self._pattern_layer, Xs)
+
+        # Layers that apply the backward pass.
+        grad_act = ilayers.GradientWRT(len(act_Xs))
+        grad_pattern = ilayers.GradientWRT(len(Xs))
+
+        # First step: propagate through the activation layer.
+        # Workaround for linear activations.
+        linear_activations = [None, keras.activations.get("linear")]
+        if self._act_layer.activation in linear_activations:
+            tmp = reversed_Ys
+        else:
+            # if linear activation this behaves strange
+            tmp = utils.to_list(grad_act(act_Xs+act_Ys+reversed_Ys))
+
+        # Second step: propagate through the pattern layer.
+        return grad_pattern(Xs+pattern_Ys+tmp)
+
+
+class PatternNet(base.OneEpochTrainerMixin, base.ReverseAnalyzerBase):
+    """PatternNet analyzer.
+
+    Applies the "PatternNet" algorithm to analyze the model's predictions.
+
+    :param model: A Keras model.
+    :param patterns: Pattern computed by
+      :class:`innvestigate.tools.PatternComputer`. If None :func:`fit` needs
+      to be called.
+    :param allow_lambda_layers: Approximate lambda layers with the gradient.
+    :param reverse_project_bottleneck_layers: Project the analysis vector into
+      range [-1, +1]. (default: True)
+    """
+
+    def __init__(self,
+                 model,
+                 patterns=None,
+                 pattern_type=None,
+                 **kwargs):
+
+        self._add_model_softmax_check()
+        self._add_model_check(
+            lambda layer: not kchecks.only_relu_activation(layer),
+            ("PatternNet is not well defined for "
+             "networks with non-ReLU activations."),
+            check_type="warning",
+        )
+        self._add_model_check(
+            lambda layer: not kchecks.is_convnet_layer(layer),
+            ("PatternNet is only well defined for "
+             "convolutional neural networks."),
+            check_type="warning",
+        )
+        self._add_model_check(
+            lambda layer: not isinstance(layer,
+                                         SUPPORTED_LAYER_PATTERNNET),
+            ("PatternNet is only well defined for "
+             "conv2d/max-pooling/dense layers."),
+            check_type="exception",
+        )
+
+        self._patterns = patterns
+        if self._patterns is not None:
+            # copy pattern references
+            self._patterns = list(patterns)
+        self._pattern_type = pattern_type
+
+        # Pattern projections can lead to +-inf value with long networks.
+        # We are only interested in the direction, therefore it is save to
+        # Prevent this by projecting the values in bottleneck layers to +-1.
+        if not kwargs.get("reverse_project_bottleneck_layers", True):
+            warnings.warn("The standard setting for "
+                          "'reverse_project_bottleneck_layers' "
+                          "is overwritten.")
+        else:
+            kwargs["reverse_project_bottleneck_layers"] = True
+
+        super(PatternNet, self).__init__(model, **kwargs)
+
+    def _get_pattern_for_layer(self, layer, state):
+        layers = [l for l in kgraph.get_model_layers(self._model)
+                  if kchecks.contains_kernel(l)]
+
+        return self._patterns[layers.index(layer)]
+
+    def _prepare_pattern(self, layer, state, pattern):
+        """""Prepares a pattern before it is set in the back-ward pass."""
+        return pattern
+
+    def _create_analysis(self, *args, **kwargs):
+
+        # Apply the pattern mapping on all layers that contain a kernel.
+        def create_kernel_layer_mapping(layer, state):
+            pattern = self._get_pattern_for_layer(layer, state)
+            pattern = self._prepare_pattern(layer, state, pattern)
+            mapping_obj = PatternNetReverseKernelLayer(layer, state, pattern)
+            return mapping_obj.apply
+        self._add_conditional_reverse_mapping(
+            kchecks.contains_kernel,
+            create_kernel_layer_mapping,
+            name="patternnet_kernel_layer_mapping"
+        )
+
+        return super(PatternNet, self)._create_analysis(*args, **kwargs)
+
+    def _fit_generator(self,
+                       generator,
+                       steps_per_epoch=None,
+                       epochs=1,
+                       max_queue_size=10,
+                       workers=1,
+                       use_multiprocessing=False,
+                       verbose=0,
+                       disable_no_training_warning=None,
+                       **kwargs):
+
+        pattern_type = self._pattern_type
+        if pattern_type is None:
+            pattern_type = "relu"
+
+        if isinstance(pattern_type, (list, tuple)):
+            raise ValueError("Only one pattern type allowed. "
+                             "Please pass a string.")
+
+        computer = itools.PatternComputer(self._model,
+                                          pattern_type=pattern_type,
+                                          **kwargs)
+
+        self._patterns = computer.compute_generator(
+            generator,
+            steps_per_epoch=steps_per_epoch,
+            max_queue_size=max_queue_size,
+            workers=workers,
+            use_multiprocessing=use_multiprocessing,
+            verbose=verbose)
+
+    def _get_state(self):
+        state = super(PatternNet, self)._get_state()
+        state.update({"patterns": self._patterns,
+                      "pattern_type": self._pattern_type})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        patterns = state.pop("patterns")
+        pattern_type = state.pop("pattern_type")
+        kwargs = super(PatternNet, clazz)._state_to_kwargs(state)
+        kwargs.update({"patterns": patterns,
+                       "pattern_type": pattern_type})
+        return kwargs
+
+
+class PatternAttribution(PatternNet):
+    """PatternAttribution analyzer.
+
+    Applies the "PatternNet" algorithm to analyze the model's predictions.
+
+    :param model: A Keras model.
+    :param patterns: Pattern computed by
+      :class:`innvestigate.tools.PatternComputer`. If None :func:`fit` needs
+      to be called.
+    :param allow_lambda_layers: Approximate lambda layers with the gradient.
+    :param reverse_project_bottleneck_layers: Project the analysis vector into
+      range [-1, +1]. (default: True)
+    """
+
+    def _prepare_pattern(self, layer, state, pattern):
+        weights = layer.get_weights()
+        tmp = [pattern.shape == x.shape for x in weights]
+        if np.sum(tmp) != 1:
+            raise Exception("Cannot match pattern to kernel.")
+        weight = weights[np.argmax(tmp)]
+        return np.multiply(pattern, weight)

+ 1 - 0
original_model/innvestigate/analyzer/relevance_based/__init__.py

@@ -0,0 +1 @@
+__all__ = ["relevance_analyzer", "relevance_rule", "utils"]

+ 869 - 0
original_model/innvestigate/analyzer/relevance_based/relevance_analyzer.py

@@ -0,0 +1,869 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import zip
+import six
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+import inspect
+import keras
+import keras.backend as K
+import keras.engine.topology
+import keras.models
+import keras.layers
+import keras.layers.convolutional
+import keras.layers.core
+import keras.layers.local
+import keras.layers.noise
+import keras.layers.normalization
+import keras.layers.pooling
+
+
+from .. import base
+from innvestigate import layers as ilayers
+from innvestigate import utils as iutils
+import innvestigate.utils.keras as kutils
+from innvestigate.utils.keras import checks as kchecks
+from innvestigate.utils.keras import graph as kgraph
+from . import relevance_rule as rrule
+from . import utils as rutils
+
+
+__all__ = [
+    "BaselineLRPZ",
+
+    "LRP",
+    "LRP_RULES",
+
+    "LRPZ",
+    "LRPZIgnoreBias",
+
+    "LRPEpsilon",
+    "LRPEpsilonIgnoreBias",
+
+    "LRPWSquare",
+    "LRPFlat",
+
+    "LRPAlphaBeta",
+
+    "LRPAlpha2Beta1",
+    "LRPAlpha2Beta1IgnoreBias",
+    "LRPAlpha1Beta0",
+    "LRPAlpha1Beta0IgnoreBias",
+    "LRPZPlus",
+    "LRPZPlusFast",
+
+    "LRPSequentialPresetA",
+    "LRPSequentialPresetB",
+
+    "LRPSequentialPresetAFlat",
+    "LRPSequentialPresetBFlat",
+    "LRPSequentialPresetBFlatUntilIdx",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class BaselineLRPZ(base.AnalyzerNetworkBase):
+    """LRPZ analyzer - for testing purpose only.
+
+    Applies the "LRP-Z" algorithm to analyze the model.
+    Based on the gradient times the input formula.
+    **This formula holds only for ReLU/MaxPooling networks, for which
+    LRP-Z collapses into the stated formula.**
+
+    :param model: A Keras model.
+    """
+
+    def __init__(self, model, **kwargs):
+        # Inside function to not break import if Keras changes.
+        BASELINELRPZ_LAYERS = (
+            keras.engine.topology.InputLayer,
+            keras.layers.convolutional.Conv1D,
+            keras.layers.convolutional.Conv2D,
+            keras.layers.convolutional.Conv2DTranspose,
+            keras.layers.convolutional.Conv3D,
+            keras.layers.convolutional.Conv3DTranspose,
+            keras.layers.convolutional.Cropping1D,
+            keras.layers.convolutional.Cropping2D,
+            keras.layers.convolutional.Cropping3D,
+            keras.layers.convolutional.SeparableConv1D,
+            keras.layers.convolutional.SeparableConv2D,
+            keras.layers.convolutional.UpSampling1D,
+            keras.layers.convolutional.UpSampling2D,
+            keras.layers.convolutional.UpSampling3D,
+            keras.layers.convolutional.ZeroPadding1D,
+            keras.layers.convolutional.ZeroPadding2D,
+            keras.layers.convolutional.ZeroPadding3D,
+            keras.layers.core.Activation,
+            keras.layers.core.ActivityRegularization,
+            keras.layers.core.Dense,
+            keras.layers.core.Dropout,
+            keras.layers.core.Flatten,
+            keras.layers.core.Lambda,
+            keras.layers.core.Masking,
+            keras.layers.core.Permute,
+            keras.layers.core.RepeatVector,
+            keras.layers.core.Reshape,
+            keras.layers.core.SpatialDropout1D,
+            keras.layers.core.SpatialDropout2D,
+            keras.layers.core.SpatialDropout3D,
+            keras.layers.local.LocallyConnected1D,
+            keras.layers.local.LocallyConnected2D,
+            keras.layers.Add,
+            keras.layers.Concatenate,
+            keras.layers.Dot,
+            keras.layers.Maximum,
+            keras.layers.Minimum,
+            keras.layers.Subtract,
+            keras.layers.noise.AlphaDropout,
+            keras.layers.noise.GaussianDropout,
+            keras.layers.noise.GaussianNoise,
+            keras.layers.normalization.BatchNormalization,
+            keras.layers.pooling.GlobalMaxPooling1D,
+            keras.layers.pooling.GlobalMaxPooling2D,
+            keras.layers.pooling.GlobalMaxPooling3D,
+            keras.layers.pooling.MaxPooling1D,
+            keras.layers.pooling.MaxPooling2D,
+            keras.layers.pooling.MaxPooling3D,
+        )
+
+        self._add_model_softmax_check()
+        self._add_model_check(
+            lambda layer: not kchecks.only_relu_activation(layer),
+            "BaselineLRPZ only works with  ReLU activations.",
+            check_type="exception",
+        )
+        self._add_model_check(
+            lambda layer: not isinstance(layer, BASELINELRPZ_LAYERS),
+            "BaselineLRPZ only works with a predefined set of layers.",
+            check_type="exception",
+        )
+
+        super(BaselineLRPZ, self).__init__(model, **kwargs)
+
+    def _create_analysis(self, model, stop_analysis_at_tensors=[]):
+        tensors_to_analyze = [x for x in iutils.to_list(model.inputs)
+                              if x not in stop_analysis_at_tensors]
+        gradients = iutils.to_list(ilayers.Gradient()(
+            tensors_to_analyze+[model.outputs[0]]))
+        return [keras.layers.Multiply()([i, g])
+                for i, g in zip(tensors_to_analyze, gradients)]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+# Utility list enabling name mappings via string
+LRP_RULES = {
+    "Z": rrule.ZRule,
+    "ZIgnoreBias": rrule.ZIgnoreBiasRule,
+
+    "Epsilon": rrule.EpsilonRule,
+    "EpsilonIgnoreBias": rrule.EpsilonIgnoreBiasRule,
+
+    "WSquare": rrule.WSquareRule,
+    "Flat": rrule.FlatRule,
+
+    "AlphaBeta": rrule.AlphaBetaRule,
+    "AlphaBetaIgnoreBias": rrule.AlphaBetaIgnoreBiasRule,
+
+    "Alpha2Beta1": rrule.Alpha2Beta1Rule,
+    "Alpha2Beta1IgnoreBias": rrule.Alpha2Beta1IgnoreBiasRule,
+    "Alpha1Beta0": rrule.Alpha1Beta0Rule,
+    "Alpha1Beta0IgnoreBias": rrule.Alpha1Beta0IgnoreBiasRule,
+
+    "ZPlus": rrule.ZPlusRule,
+    "ZPlusFast": rrule.ZPlusFastRule,
+    "Bounded": rrule.BoundedRule,
+}
+
+
+class EmbeddingReverseLayer(kgraph.ReverseMappingBase):
+    def __init__(self, layer, state):
+        #TODO: implement rule support.
+        return
+
+    def apply(self, Xs, Ys, Rs, reverse_state):
+        # the embedding layer outputs for an (indexed) input a vector.
+        # thus, in the relevance backward pass, the embedding layer receives
+        # relevances Rs corresponding to those vectors.
+        # Due to the 1:1 relationship between input index and output mapping vector,
+        # the relevance backward pass can be realized by pooling relevances
+        # over the vector axis.
+
+        #relevances are given shaped [batch_size, sequence_length, embedding_dims]
+        pool_relevance = keras.layers.Lambda(lambda x: keras.backend.sum(x, axis=-1))
+        return [pool_relevance(r) for r in Rs]
+
+class BatchNormalizationReverseLayer(kgraph.ReverseMappingBase):
+    """Special BN handler that applies the Z-Rule"""
+
+    def __init__(self, layer, state):
+        config = layer.get_config()
+
+        self._center = config['center']
+        self._scale = config['scale']
+        self._axis = config['axis']
+
+        self._mean = layer.moving_mean
+        self._std = layer.moving_variance
+        if self._center:
+            self._beta = layer.beta
+
+        #TODO: implement rule support. for BatchNormalization -> [BNEpsilon, BNAlphaBeta, BNIgnore]
+        #super(BatchNormalizationReverseLayer, self).__init__(layer, state)
+        # how to do this:
+        # super.__init__ calls select_rule and sets a self._rule class
+        # check if isinstance(self_rule, EpsiloneRule), then reroute
+        # to BatchNormEpsilonRule. Not pretty, but should work.
+
+    def apply(self, Xs, Ys, Rs, reverse_state):
+        input_shape = [K.int_shape(x) for x in Xs]
+        if len(input_shape) != 1:
+            #extend below lambda layers towards multiple parameters.
+            raise ValueError("BatchNormalizationReverseLayer expects Xs with len(Xs) = 1, but was len(Xs) = {}".format(len(Xs)))
+        input_shape = input_shape[0]
+
+        # prepare broadcasting shape for layer parameters
+        broadcast_shape = [1] * len(input_shape)
+        broadcast_shape[self._axis] = input_shape[self._axis]
+        broadcast_shape[0] =  -1
+
+        #reweight relevances as
+        #        x * (y - beta)     R
+        # Rin = ---------------- * ----
+        #           x - mu          y
+        # batch norm can be considered as 3 distinct layers of subtraction,
+        # multiplication and then addition. The multiplicative scaling layer
+        # has no effect on LRP and functions as a linear activation layer
+
+        minus_mu = keras.layers.Lambda(lambda x: x - K.reshape(self._mean, broadcast_shape))
+        minus_beta = keras.layers.Lambda(lambda x: x - K.reshape(self._beta, broadcast_shape))
+        prepare_div = keras.layers.Lambda(lambda x: x + (K.cast(K.greater_equal(x,0), K.floatx())*2-1)*K.epsilon())
+
+
+        x_minus_mu = kutils.apply(minus_mu, Xs)
+        if self._center:
+            y_minus_beta = kutils.apply(minus_beta, Ys)
+        else:
+            y_minus_beta = Ys
+
+        numerator = [keras.layers.Multiply()([x, ymb, r])
+                     for x, ymb, r in zip(Xs, y_minus_beta, Rs)]
+        denominator = [keras.layers.Multiply()([xmm, y])
+                       for xmm, y in zip(x_minus_mu, Ys)]
+
+        return [ilayers.SafeDivide()([n, prepare_div(d)])
+                for n, d in zip(numerator, denominator)]
+
+
+class AddReverseLayer(kgraph.ReverseMappingBase):
+    """Special Add layer handler that applies the Z-Rule"""
+
+    def __init__(self, layer, state):
+        self._layer_wo_act = kgraph.copy_layer_wo_activation(layer,
+                                                             name_template="reversed_kernel_%s")
+
+        #TODO: implement rule support.
+        #super(AddReverseLayer, self).__init__(layer, state)
+
+    def apply(self, Xs, Ys, Rs, reverse_state):
+        # the outputs of the pooling operation at each location is the sum of its inputs.
+        # the forward message must be known in this case, and are the inputs for each pooling thing.
+        # the gradient is 1 for each output-to-input connection, which corresponds to the "weights"
+        # of the layer. It should thus be sufficient to reweight the relevances and and do a gradient_wrt
+        grad = ilayers.GradientWRT(len(Xs))
+        # Get activations.
+        Zs = kutils.apply(self._layer_wo_act, Xs)
+        # Divide incoming relevance by the activations.
+        tmp = [ilayers.SafeDivide()([a, b])
+               for a, b in zip(Rs, Zs)]
+
+        # Propagate the relevance to input neurons
+        # using the gradient.
+        tmp = iutils.to_list(grad(Xs+Zs+tmp))
+        # Re-weight relevance with the input values.
+        return [keras.layers.Multiply()([a, b])
+                for a, b in zip(Xs, tmp)]
+
+
+class AveragePoolingReverseLayer(kgraph.ReverseMappingBase):
+    """Special AveragePooling handler that applies the Z-Rule"""
+
+    def __init__(self, layer, state):
+        self._layer_wo_act = kgraph.copy_layer_wo_activation(layer,
+                                                             name_template="reversed_kernel_%s")
+
+        #TODO: implement rule support.
+        #super(AveragePoolingRerseLayer, self).__init__(layer, state)
+
+    def apply(self, Xs, Ys, Rs, reverse_state):
+        # the outputs of the pooling operation at each location is the sum of its inputs.
+        # the forward message must be known in this case, and are the inputs for each pooling thing.
+        # the gradient is 1 for each output-to-input connection, which corresponds to the "weights"
+        # of the layer. It should thus be sufficient to reweight the relevances and and do a gradient_wrt
+
+        grad = ilayers.GradientWRT(len(Xs))
+        # Get activations.
+        Zs = kutils.apply(self._layer_wo_act, Xs)
+        # Divide incoming relevance by the activations.
+        tmp = [ilayers.SafeDivide()([a, b])
+               for a, b in zip(Rs, Zs)]
+
+        # Propagate the relevance to input neurons
+        # using the gradient.
+        tmp = iutils.to_list(grad(Xs+Zs+tmp))
+        # Re-weight relevance with the input values.
+        return [keras.layers.Multiply()([a, b])
+                for a, b in zip(Xs, tmp)]
+
+
+class LRP(base.ReverseAnalyzerBase):
+    """
+    Base class for LRP-based model analyzers
+
+
+    :param model: A Keras model.
+
+    :param rule: A rule can be a  string or a Rule object, lists thereof or a list of conditions [(Condition, Rule), ... ]
+      gradient.
+
+    :param input_layer_rule: either a Rule object, atuple of (low, high) the min/max pixel values of the inputs
+    :param bn_layer_rule: either a Rule object or None.
+      None means dedicated BN rule will be applied.
+    """
+
+    def __init__(self, model, *args, **kwargs):
+        rule = kwargs.pop("rule", None)
+        input_layer_rule = kwargs.pop("input_layer_rule", None)
+        until_layer_idx = kwargs.pop("until_layer_idx", None)
+        until_layer_rule = kwargs.pop("until_layer_rule", None)
+
+        bn_layer_rule = kwargs.pop("bn_layer_rule", None)
+        bn_layer_fuse_mode = kwargs.pop("bn_layer_fuse_mode", "one_linear")
+        assert bn_layer_fuse_mode in ["one_linear", "two_linear"]
+
+        self._add_model_softmax_check()
+        self._add_model_check(
+            lambda layer: not kchecks.is_convnet_layer(layer),
+            "LRP is only tested for convolutional neural networks.",
+            check_type="warning",
+        )
+
+        # check if rule was given explicitly.
+        # rule can be a string, a list (of strings) or a list of conditions [(Condition, Rule), ... ] for each layer.
+        if rule is None:
+            raise ValueError("Need LRP rule(s).")
+
+
+
+        if isinstance(rule, list):
+            # copy refrences
+            self._rule = list(rule)
+        else:
+            self._rule = rule
+        self._input_layer_rule = input_layer_rule
+        self._until_layer_rule = until_layer_rule
+        self._until_layer_idx = until_layer_idx
+
+        self._bn_layer_rule = bn_layer_rule
+        self._bn_layer_fuse_mode = bn_layer_fuse_mode
+
+        if(
+           isinstance(rule, six.string_types) or
+           (inspect.isclass(rule) and issubclass(rule, kgraph.ReverseMappingBase)) # NOTE: All LRP rules inherit from kgraph.ReverseMappingBase
+        ):
+            # the given rule is a single string or single rule implementing cla ss
+            use_conditions = True
+            rules = [(lambda a, b: True, rule)]
+
+        elif not isinstance(rule[0], tuple):
+            # rule list of rule strings or classes
+            use_conditions = False
+            rules = list(rule)
+        else:
+            # rule is list of conditioned rules
+            use_conditions = True
+            rules = rule
+
+        #apply rule to first self._until_layer_idx layers
+        if self._until_layer_rule is not None and self._until_layer_idx is not None:
+            for i in range(self._until_layer_idx+1):
+                rules.insert(0,
+                             (lambda layer, foo, bound_i=i: kchecks.is_layer_at_idx(layer, bound_i),
+                              self._until_layer_rule))
+
+        # create a BoundedRule for input layer handling from given tuple
+        if self._input_layer_rule is not None:
+            input_layer_rule = self._input_layer_rule
+            if isinstance(input_layer_rule, tuple):
+                low, high = input_layer_rule
+
+                class BoundedProxyRule(rrule.BoundedRule):
+                    def __init__(self, *args, **kwargs):
+                        super(BoundedProxyRule, self).__init__(
+                            *args, low=low, high=high, **kwargs)
+                input_layer_rule = BoundedProxyRule
+
+
+            if use_conditions is True:
+                rules.insert(0,
+                             (lambda layer, foo: kchecks.is_input_layer(layer),
+                              input_layer_rule))
+
+            else:
+                rules.insert(0, input_layer_rule)
+
+        self._rules_use_conditions = use_conditions
+        self._rules = rules
+
+        # FINALIZED constructor.
+        super(LRP, self).__init__(model, *args, **kwargs)
+
+    def create_rule_mapping(self, layer, reverse_state):
+        rule_class = None
+        if self._rules_use_conditions is True:
+            for condition, rule in self._rules:
+                if condition(layer, reverse_state):
+                    rule_class = rule
+                    break
+        else:
+            rule_class = self._rules.pop()
+
+        if rule_class is None:
+            raise Exception("No rule applies to layer: %s" % layer)
+
+        if isinstance(rule_class, six.string_types):
+            rule_class = LRP_RULES[rule_class]
+        rule = rule_class(layer, reverse_state)
+
+        return rule.apply
+
+    def _create_analysis(self, *args, **kwargs):
+        ####################################################################
+        ### Functionality responible for backwards rule selection below ####
+        ####################################################################
+
+        # default backward hook
+        self._add_conditional_reverse_mapping(
+            kchecks.contains_kernel,
+            self.create_rule_mapping,
+            name="lrp_layer_with_kernel_mapping",
+        )
+
+        #specialized backward hooks. TODO: add ReverseLayer class handling layers Without kernel: Add and AvgPool
+        bn_layer_rule = self._bn_layer_rule
+
+        if bn_layer_rule is None:
+            # todo(alber): get rid of this option!
+            # alternatively a default rule should be applied.
+            bn_mapping = BatchNormalizationReverseLayer
+        else:
+            if isinstance(bn_layer_rule, six.string_types):
+                bn_layer_rule = LRP_RULES[bn_layer_rule]
+
+            bn_mapping = kgraph.apply_mapping_to_fused_bn_layer(
+                bn_layer_rule,
+                fuse_mode=self._bn_layer_fuse_mode,
+            )
+        self._add_conditional_reverse_mapping(
+            kchecks.is_batch_normalization_layer,
+            bn_mapping,
+            name="lrp_batch_norm_mapping",
+        )
+        self._add_conditional_reverse_mapping(
+            kchecks.is_average_pooling,
+            AveragePoolingReverseLayer,
+            name="lrp_average_pooling_mapping",
+        )
+        self._add_conditional_reverse_mapping(
+            kchecks.is_add_layer,
+            AddReverseLayer,
+            name="lrp_add_layer_mapping",
+        )
+        self._add_conditional_reverse_mapping(
+            kchecks.is_embedding_layer,
+            EmbeddingReverseLayer,
+            name="lrp_embedding_mapping"
+        )
+
+        # FINALIZED constructor.
+        return super(LRP, self)._create_analysis(*args, **kwargs)
+
+
+    def _default_reverse_mapping(self, Xs, Ys, reversed_Ys, reverse_state):
+        #default_return_layers = [keras.layers.Activation]# TODO extend
+        if(len(Xs) == len(Ys) and
+           isinstance(reverse_state['layer'], (keras.layers.Activation,)) and
+           all([K.int_shape(x) == K.int_shape(y) for x, y in zip(Xs, Ys)])):
+            # Expect Xs and Ys to have the same shapes.
+            # There is not mixing of relevances as there is kernel,
+            # therefore we pass them as they are.
+            return reversed_Ys
+        else:
+            # This branch covers:
+            # MaxPooling
+            # Max
+            # Flatten
+            # Reshape
+            # Concatenate
+            # Cropping
+            return self._gradient_reverse_mapping(
+                Xs, Ys, reversed_Ys, reverse_state)
+
+    ########################################
+    ### End of Rule Selection Business. ####
+    ########################################
+
+
+    def _get_state(self):
+        state = super(LRP, self)._get_state()
+        state.update({"rule": self._rule})
+        state.update({"input_layer_rule": self._input_layer_rule})
+        state.update({"bn_layer_rule": self._bn_layer_rule})
+        state.update({"bn_layer_fuse_mode": self._bn_layer_fuse_mode})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        rule = state.pop("rule")
+        input_layer_rule = state.pop("input_layer_rule")
+        bn_layer_rule = state.pop("bn_layer_rule")
+        bn_layer_fuse_mode = state.pop("bn_layer_fuse_mode")
+        kwargs = super(LRP, clazz)._state_to_kwargs(state)
+        kwargs.update({"rule": rule,
+                       "input_layer_rule": input_layer_rule,
+                       "bn_layer_rule": bn_layer_rule,
+                       "bn_layer_fuse_mode": bn_layer_fuse_mode})
+        return kwargs
+
+
+###############################################################################
+# ANALYZER CLASSES AND PRESETS ################################################
+###############################################################################
+
+
+class _LRPFixedParams(LRP):
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        kwargs = super(_LRPFixedParams, clazz)._state_to_kwargs(state)
+        del kwargs["rule"]
+        del kwargs["bn_layer_rule"]
+        return kwargs
+
+
+class LRPZ(_LRPFixedParams):
+    """LRP-analyzer that uses the LRP-Z rule"""
+    
+    def __init__(self, model, *args, **kwargs):
+        super(LRPZ, self).__init__(model, *args,
+                                   rule="Z", bn_layer_rule="Z", **kwargs)
+
+
+class LRPZIgnoreBias(_LRPFixedParams):
+    """LRP-analyzer that uses the LRP-Z-ignore-bias rule"""
+
+    def __init__(self, model, *args, **kwargs):
+        super(LRPZIgnoreBias, self).__init__(model, *args,
+                                             rule="ZIgnoreBias",
+                                             bn_layer_rule="ZIgnoreBias",
+                                             **kwargs)
+
+
+
+class LRPEpsilon(_LRPFixedParams):
+    """LRP-analyzer that uses the LRP-Epsilon rule"""
+
+    def __init__(self, model, epsilon=1e-7, bias=True, *args, **kwargs):
+        epsilon = rutils.assert_lrp_epsilon_param(epsilon, self)
+        self._epsilon = epsilon
+
+        class EpsilonProxyRule(rrule.EpsilonRule):
+            """
+            Dummy class inheriting from EpsilonRule
+            for passing along the chosen parameters from
+            the LRP analyzer class to the decopmosition rules.
+            """
+            def __init__(self, *args, **kwargs):
+                super(EpsilonProxyRule, self).__init__(*args,
+                                                       epsilon=epsilon,
+                                                       bias=bias,
+                                                       **kwargs)
+
+        super(LRPEpsilon, self).__init__(model, *args,
+                                         rule=EpsilonProxyRule,
+                                         bn_layer_rule=EpsilonProxyRule,
+                                         **kwargs)
+
+
+class LRPEpsilonIgnoreBias(LRPEpsilon):
+    """LRP-analyzer that uses the LRP-Epsilon-ignore-bias rule"""
+
+    def __init__(self, model, epsilon=1e-7, *args, **kwargs):
+        super(LRPEpsilonIgnoreBias, self).__init__(model, *args,
+                                                   epsilon=epsilon,
+                                                   bias=False,
+                                                   **kwargs)
+
+
+class LRPWSquare(_LRPFixedParams):
+    """LRP-analyzer that uses the DeepTaylor W**2 rule"""
+
+    def __init__(self, model, *args, **kwargs):
+        super(LRPWSquare, self).__init__(model, *args,
+                                         rule="WSquare",
+                                         bn_layer_rule="WSquare",
+                                         **kwargs)
+
+
+class LRPFlat(_LRPFixedParams):
+    """LRP-analyzer that uses the LRP-Flat rule"""
+
+    def __init__(self, model, *args, **kwargs):
+        super(LRPFlat, self).__init__(model, *args,
+                                      rule="Flat",
+                                      bn_layer_rule="Flat",
+                                      **kwargs)
+
+
+class LRPAlphaBeta(LRP):
+    """ Base class for LRP AlphaBeta"""
+
+    def __init__(self, model, alpha=None, beta=None, bias=True, *args, **kwargs):
+        alpha, beta = rutils.assert_infer_lrp_alpha_beta_param(alpha, beta, self)
+        self._alpha = alpha
+        self._beta = beta
+        self._bias = bias
+
+        class AlphaBetaProxyRule(rrule.AlphaBetaRule):
+            """
+            Dummy class inheriting from AlphaBetaRule
+            for the purpose of passing along the chosen parameters from
+            the LRP analyzer class to the decopmosition rules.
+            """
+            def __init__(self, *args, **kwargs):
+                super(AlphaBetaProxyRule, self).__init__(*args,
+                                                         alpha=alpha,
+                                                         beta=beta,
+                                                         bias=bias,
+                                                         **kwargs)
+
+        super(LRPAlphaBeta, self).__init__(model, *args,
+                                           rule=AlphaBetaProxyRule,
+                                           bn_layer_rule=AlphaBetaProxyRule,
+                                           **kwargs)
+
+    def _get_state(self):
+        state = super(LRPAlphaBeta, self)._get_state()
+        del state["rule"]
+        state.update({"alpha": self._alpha})
+        state.update({"beta": self._beta})
+        state.update({"bias": self._bias})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        alpha = state.pop("alpha")
+        beta = state.pop("beta")
+        bias = state.pop("bias")
+        state["rule"] = None
+        kwargs = super(LRPAlphaBeta, clazz)._state_to_kwargs(state)
+        del kwargs["rule"]
+        del kwargs["bn_layer_rule"]
+        kwargs.update({"alpha": alpha,
+                       "beta": beta,
+                       "bias": bias})
+        return kwargs
+
+
+
+
+
+
+class _LRPAlphaBetaFixedParams(LRPAlphaBeta):
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        kwargs = super(_LRPAlphaBetaFixedParams, clazz)._state_to_kwargs(state)
+        del kwargs["alpha"]
+        del kwargs["beta"]
+        del kwargs["bias"]
+        return kwargs
+
+
+class LRPAlpha2Beta1(_LRPAlphaBetaFixedParams):
+    """LRP-analyzer that uses the LRP-alpha-beta rule with a=2,b=1"""
+
+    def __init__(self, model, *args, **kwargs):
+        super(LRPAlpha2Beta1, self).__init__(model, *args,
+                                             alpha=2,
+                                             beta=1,
+                                             bias=True,
+                                             **kwargs)
+
+
+class LRPAlpha2Beta1IgnoreBias(_LRPAlphaBetaFixedParams):
+    """LRP-analyzer that uses the LRP-alpha-beta-ignbias rule with a=2,b=1"""
+
+    def __init__(self, model, *args, **kwargs):
+        super(LRPAlpha2Beta1IgnoreBias, self).__init__(model, *args,
+                                                       alpha=2,
+                                                       beta=1,
+                                                       bias=False,
+                                                       **kwargs)
+
+
+class LRPAlpha1Beta0(_LRPAlphaBetaFixedParams):
+    """LRP-analyzer that uses the LRP-alpha-beta rule with a=1,b=0"""
+
+    def __init__(self, model, *args, **kwargs):
+        super(LRPAlpha1Beta0, self).__init__(model, *args,
+                                             alpha=1,
+                                             beta=0,
+                                             bias=True,
+                                             **kwargs)
+
+
+class LRPAlpha1Beta0IgnoreBias(_LRPAlphaBetaFixedParams):
+    """LRP-analyzer that uses the LRP-alpha-beta-ignbias rule with a=1,b=0"""
+
+    def __init__(self, model, *args, **kwargs):
+        super(LRPAlpha1Beta0IgnoreBias, self).__init__(model, *args,
+                                                       alpha=1,
+                                                       beta=0,
+                                                       bias=False,
+                                                       **kwargs)
+
+class LRPZPlus(LRPAlpha1Beta0IgnoreBias):
+    """LRP-analyzer that uses the LRP-alpha-beta rule with a=1,b=0"""
+    #TODO: assert that layer inputs are always >= 0
+    def __init__(self, model, *args, **kwargs):
+        super(LRPZPlus, self).__init__(model, *args, **kwargs)
+
+
+class LRPZPlusFast(_LRPFixedParams):
+    """
+    The ZPlus rule is a special case of the AlphaBetaRule
+    for alpha=1, beta=0 and assumes inputs x >= 0.
+    """
+    #TODO: assert that layer inputs are always >= 0
+    def __init__(self, model, *args, **kwargs):
+        super(LRPZPlusFast, self).__init__(model, *args,
+                                           rule="ZPlusFast",
+                                           bn_layer_rule="ZPlusFast",
+                                           **kwargs)
+
+
+class LRPSequentialPresetA(_LRPFixedParams): #for the lack of a better name
+    """Special LRP-configuration for ConvNets"""
+
+    def __init__(self, model, epsilon=1e-1, *args, **kwargs):
+
+        self._add_model_check(
+            lambda layer: not kchecks.only_relu_activation(layer),
+            #TODO: fix. specify. extend.
+            ("LRPSequentialPresetA is not advised "
+             "for networks with non-ReLU activations."),
+            check_type="warning",
+        )
+
+        class EpsilonProxyRule(rrule.EpsilonRule):
+            def __init__(self, *args, **kwargs):
+                super(EpsilonProxyRule, self).__init__(*args,
+                                                       epsilon=epsilon,
+                                                       bias=True,
+                                                       **kwargs)
+
+
+        conditional_rules = [(kchecks.is_dense_layer, EpsilonProxyRule),
+                             (kchecks.is_conv_layer, rrule.Alpha1Beta0Rule)
+                            ]
+        bn_layer_rule = kwargs.pop("bn_layer_rule", rrule.AlphaBetaX2m100Rule)
+
+        super(LRPSequentialPresetA, self).__init__(
+            model,
+            *args,
+            rule=conditional_rules,
+            bn_layer_rule=bn_layer_rule,
+            **kwargs)
+
+
+class LRPSequentialPresetB(_LRPFixedParams):
+    """Special LRP-configuration for ConvNets"""
+
+    def __init__(self, model, epsilon=1e-1, *args, **kwargs):
+        self._add_model_check(
+            lambda layer: not kchecks.only_relu_activation(layer),
+            #TODO: fix. specify. extend.
+            ("LRPSequentialPresetB is not advised "
+             "for networks with non-ReLU activations."),
+            check_type="warning",
+        )
+
+        class EpsilonProxyRule(rrule.EpsilonRule):
+            def __init__(self, *args, **kwargs):
+                super(EpsilonProxyRule, self).__init__(*args,
+                                                       epsilon=epsilon,
+                                                       bias=True,
+                                                       **kwargs)
+
+        conditional_rules = [(kchecks.is_dense_layer, EpsilonProxyRule),
+                             (kchecks.is_conv_layer, rrule.Alpha2Beta1Rule)
+                         ]
+        bn_layer_rule = kwargs.pop("bn_layer_rule", rrule.AlphaBetaX2m100Rule)
+
+        super(LRPSequentialPresetB, self).__init__(
+            model,
+            *args,
+            rule=conditional_rules,
+            bn_layer_rule=bn_layer_rule,
+            **kwargs)
+
+
+
+
+
+#TODO: allow to pass input layer identification by index or id.
+class LRPSequentialPresetAFlat(LRPSequentialPresetA):
+    """Special LRP-configuration for ConvNets"""
+
+    def __init__(self, model, *args, **kwargs):
+        super(LRPSequentialPresetAFlat, self).__init__(model,
+                                                *args,
+                                                input_layer_rule="Flat",
+                                                **kwargs)
+
+
+
+#TODO: allow to pass input layer identification by index or id.
+class LRPSequentialPresetBFlat(LRPSequentialPresetB):
+    """Special LRP-configuration for ConvNets"""
+
+    def __init__(self, model, *args, **kwargs):
+        super(LRPSequentialPresetBFlat, self).__init__(model,
+                                                *args,
+                                                input_layer_rule="Flat",
+                                                **kwargs)
+
+class LRPSequentialPresetBFlatUntilIdx(LRPSequentialPresetBFlat):
+    """
+        Special LRP-configuration for ConvNets. Allows to perform LRP_flat from (including) layer until_layer_idx down until
+        the input layer. Weightless layers are ignored when counting the index for now.
+    """
+
+    def __init__(self, model, *args, **kwargs):
+        layer_flat_idx=kwargs.pop("until_layer_idx", None)
+        super(LRPSequentialPresetBFlatUntilIdx, self).__init__(model,
+                                                    *args,
+                                                    until_layer_idx=layer_flat_idx,
+                                                    until_layer_rule=rrule.FlatRule,
+                                                    **kwargs)

+ 635 - 0
original_model/innvestigate/analyzer/relevance_based/relevance_rule.py

@@ -0,0 +1,635 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import zip
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+import keras
+import keras.backend as K
+import keras.engine.topology
+import keras.models
+import keras.layers
+import keras.layers.convolutional
+import keras.layers.core
+import keras.layers.local
+import keras.layers.noise
+import keras.layers.normalization
+import keras.layers.pooling
+import numpy as np
+
+
+from innvestigate import layers as ilayers
+from innvestigate import utils as iutils
+import innvestigate.utils.keras as kutils
+from innvestigate.utils.keras import backend as iK
+from innvestigate.utils.keras import graph as kgraph
+from . import utils as rutils
+
+
+# TODO: differentiate between LRP and DTD rules?
+# DTD rules are special cases of LRP rules with additional assumptions
+__all__ = [
+    #dedicated treatment for special layers
+
+
+    #general rules
+    "ZRule",
+    "ZIgnoreBiasRule",
+
+    "EpsilonRule",
+    "EpsilonIgnoreBiasRule",
+
+    "WSquareRule",
+    "FlatRule",
+
+    "AlphaBetaRule",
+    "AlphaBetaIgnoreBiasRule",
+
+    "Alpha2Beta1Rule",
+    "Alpha2Beta1IgnoreBiasRule",
+
+    "Alpha1Beta0Rule",
+    "Alpha1Beta0IgnoreBiasRule",
+
+    "AlphaBetaXRule",
+    "AlphaBetaX1000Rule",
+    "AlphaBetaX1010Rule",
+    "AlphaBetaX1001Rule",
+    "AlphaBetaX2m100Rule",
+
+    "ZPlusRule",
+    "ZPlusFastRule",
+    "BoundedRule"
+]
+
+
+
+class ZRule(kgraph.ReverseMappingBase):
+    """
+    Basic LRP decomposition rule (for layers with weight kernels),
+    which considers the bias a constant input neuron.
+    """
+
+    def __init__(self, layer, state, bias=True):
+        self._layer_wo_act = kgraph.copy_layer_wo_activation(layer,
+                                                             keep_bias=bias,
+                                                             name_template="reversed_kernel_%s")
+
+    def apply(self, Xs, Ys, Rs, reverse_state):
+        grad = ilayers.GradientWRT(len(Xs))
+
+        # Get activations.
+        Zs = kutils.apply(self._layer_wo_act, Xs)
+        # Divide incoming relevance by the activations.
+        tmp = [ilayers.SafeDivide()([a, b])
+               for a, b in zip(Rs, Zs)]
+        # Propagate the relevance to input neurons
+        # using the gradient.
+        tmp = iutils.to_list(grad(Xs+Zs+tmp))
+        # Re-weight relevance with the input values.
+        return [keras.layers.Multiply()([a, b])
+                for a, b in zip(Xs, tmp)]
+
+
+
+class ZIgnoreBiasRule(ZRule):
+    """
+    Basic LRP decomposition rule, ignoring the bias neuron
+    """
+    def __init__(self, *args, **kwargs):
+        super(ZIgnoreBiasRule, self).__init__(*args,
+                                              bias=False,
+                                              **kwargs)
+
+
+
+class EpsilonRule(kgraph.ReverseMappingBase):
+    """
+    Similar to ZRule.
+    The only difference is the addition of a numerical stabilizer term
+    epsilon to the decomposition function's denominator.
+    the sign of epsilon depends on the sign of the output activation
+    0 is considered to be positive, ie sign(0) = 1
+    """
+
+    def __init__(self, layer, state, epsilon = 1e-7, bias=True):
+        self._epsilon = rutils.assert_lrp_epsilon_param(epsilon, self)
+        self._layer_wo_act = kgraph.copy_layer_wo_activation(
+            layer, keep_bias=bias, name_template="reversed_kernel_%s")
+
+
+    def apply(self, Xs, Ys, Rs, reverse_state):
+        grad = ilayers.GradientWRT(len(Xs))
+        # The epsilon rule aligns epsilon with the (extended) sign: 0 is considered to be positive
+        prepare_div = keras.layers.Lambda(lambda x: x + (K.cast(K.greater_equal(x,0), K.floatx())*2-1)*self._epsilon)
+
+        # Get activations.
+        Zs = kutils.apply(self._layer_wo_act, Xs)
+
+        # Divide incoming relevance by the activations.
+        tmp = [ilayers.Divide()([a, prepare_div(b)])
+               for a, b in zip(Rs, Zs)]
+        # Propagate the relevance to input neurons
+        # using the gradient.
+        tmp = iutils.to_list(grad(Xs+Zs+tmp))
+        # Re-weight relevance with the input values.
+        return [keras.layers.Multiply()([a, b])
+                for a, b in zip(Xs, tmp)]
+
+
+
+class EpsilonIgnoreBiasRule(EpsilonRule):
+    """Same as EpsilonRule but ignores the bias."""
+    def __init__(self, *args, **kwargs):
+        super(EpsilonIgnoreBiasRule, self).__init__(*args,
+                                                    bias=False,
+                                                    **kwargs)
+
+
+
+class WSquareRule(kgraph.ReverseMappingBase):
+    """W**2 rule from Deep Taylor Decomposition"""
+
+    def __init__(self, layer, state, copy_weights=False):
+        # W-square rule works with squared weights and no biases.
+        if copy_weights:
+            weights = layer.get_weights()
+        else:
+            weights = layer.weights
+        if layer.use_bias:
+            weights = weights[:-1]
+        weights = [x**2 for x in weights]
+
+        self._layer_wo_act_b = kgraph.copy_layer_wo_activation(
+            layer,
+            keep_bias=False,
+            weights=weights,
+            name_template="reversed_kernel_%s")
+
+
+    def apply(self, Xs, Ys, Rs, reverse_state):
+        grad = ilayers.GradientWRT(len(Xs))
+        # Create dummy forward path to take the derivative below.
+        Ys = kutils.apply(self._layer_wo_act_b, Xs)
+
+        # Compute the sum of the weights.
+        ones = ilayers.OnesLike()(Xs)
+        Zs = iutils.to_list(self._layer_wo_act_b(ones))
+        # Weight the incoming relevance.
+        tmp = [ilayers.SafeDivide()([a, b])
+               for a, b in zip(Rs, Zs)]
+        # Redistribute the relevances along the gradient.
+        tmp = iutils.to_list(grad(Xs+Ys+tmp))
+        return tmp
+
+
+
+
+class FlatRule(WSquareRule):
+    """Same as W**2 rule but sets all weights to ones."""
+
+    def __init__(self, layer, state, copy_weights=False):
+        # The flat rule works with weights equal to one and
+        # no biases.
+        if copy_weights:
+            weights = layer.get_weights()
+            if layer.use_bias:
+                weights = weights[:-1]
+            weights = [np.ones_like(x) for x in weights]
+        else:
+            weights = layer.weights
+            if layer.use_bias:
+                weights = weights[:-1]
+            weights = [K.ones_like(x) for x in weights]
+
+        self._layer_wo_act_b = kgraph.copy_layer_wo_activation(
+            layer,
+            keep_bias=False,
+            weights=weights,
+            name_template="reversed_kernel_%s")
+
+
+
+
+class AlphaBetaRule(kgraph.ReverseMappingBase):
+    """
+    This decomposition rule handles the positive forward
+    activations (x*w > 0) and negative forward activations
+    (w * x < 0) independently, reducing the risk of zero
+    divisions considerably. In fact, the only case where
+    divisions by zero can happen is if there are either
+    no positive or no negative parts to the activation
+    at all.
+    Corresponding parameterization of this rule implement
+    methods such as Excitation Backpropagation with
+    alpha=1, beta=0
+    s.t.
+    alpha - beta = 1 (after current param. scheme.)
+    and
+    alpha > 1
+    beta > 0
+    """
+
+
+    def __init__(self,
+                 layer,
+                 state,
+                 alpha=None,
+                 beta=None,
+                 bias=True,
+                 copy_weights=False):
+        alpha, beta = rutils.assert_infer_lrp_alpha_beta_param(alpha, beta, self)
+        self._alpha = alpha
+        self._beta = beta
+
+        # prepare positive and negative weights for computing positive
+        # and negative preactivations z in apply_accordingly.
+        if copy_weights:
+            weights = layer.get_weights()
+            if not bias and layer.use_bias:
+                weights = weights[:-1]
+            positive_weights = [x * (x > 0) for x in weights]
+            negative_weights = [x * (x < 0) for x in weights]
+        else:
+            weights = layer.weights
+            if not bias and layer.use_bias:
+                weights = weights[:-1]
+            positive_weights = [x * iK.to_floatx(x > 0) for x in weights]
+            negative_weights = [x * iK.to_floatx(x < 0) for x in weights]
+
+        self._layer_wo_act_positive = kgraph.copy_layer_wo_activation(
+            layer,
+            keep_bias=bias,
+            weights=positive_weights,
+            name_template="reversed_kernel_positive_%s")
+        self._layer_wo_act_negative = kgraph.copy_layer_wo_activation(
+            layer,
+            keep_bias=bias,
+            weights=negative_weights,
+            name_template="reversed_kernel_negative_%s")
+
+
+    def apply(self, Xs, Ys, Rs, reverse_state):
+        #this method is correct, but wasteful
+        grad = ilayers.GradientWRT(len(Xs))
+        times_alpha = keras.layers.Lambda(lambda x: x * self._alpha)
+        times_beta = keras.layers.Lambda(lambda x: x * self._beta)
+        keep_positives = keras.layers.Lambda(lambda x: x * K.cast(K.greater(x,0), K.floatx()))
+        keep_negatives = keras.layers.Lambda(lambda x: x * K.cast(K.less(x,0), K.floatx()))
+
+
+        def f(layer1, layer2, X1, X2):
+            # Get activations of full positive or negative part.
+            Z1 = kutils.apply(layer1, X1)
+            Z2 = kutils.apply(layer2, X2)
+            Zs = [keras.layers.Add()([a, b])
+                    for a, b in zip(Z1, Z2)]
+            # Divide incoming relevance by the activations.
+            tmp = [ilayers.SafeDivide()([a, b])
+                    for a, b in zip(Rs, Zs)]
+            # Propagate the relevance to the input neurons
+            # using the gradient
+            tmp1 = iutils.to_list(grad(X1+Z1+tmp))
+            tmp2 = iutils.to_list(grad(X2+Z2+tmp))
+            # Re-weight relevance with the input values.
+            tmp1 = [keras.layers.Multiply()([a, b])
+                    for a, b in zip(X1, tmp1)]
+            tmp2 = [keras.layers.Multiply()([a, b])
+                    for a, b in zip(X2, tmp2)]
+            #combine and return
+            return [keras.layers.Add()([a, b])
+                    for a, b in zip(tmp1, tmp2)]
+
+
+        # Distinguish postive and negative inputs.
+        Xs_pos = kutils.apply(keep_positives, Xs)
+        Xs_neg = kutils.apply(keep_negatives, Xs)
+        # xpos*wpos + xneg*wneg
+        activator_relevances = f(self._layer_wo_act_positive,
+                                 self._layer_wo_act_negative,
+                                 Xs_pos, Xs_neg)
+
+        if self._beta: #only compute beta-weighted contributions of beta is not zero
+            # xpos*wneg + xneg*wpos
+            inhibitor_relevances = f(self._layer_wo_act_negative,
+                                     self._layer_wo_act_positive,
+                                     Xs_pos, Xs_neg)
+            return [keras.layers.Subtract()([times_alpha(a), times_beta(b)])
+                        for a, b in zip(activator_relevances, inhibitor_relevances)]
+        else:
+            return activator_relevances
+
+
+
+        
+class AlphaBetaIgnoreBiasRule(AlphaBetaRule):
+    """Same as AlphaBetaRule but ignores biases."""
+    def __init__(self, *args, **kwargs):
+        super(AlphaBetaIgnoreBiasRule, self).__init__(*args,
+                                                      bias=False,
+                                                      **kwargs)
+
+
+
+class Alpha2Beta1Rule(AlphaBetaRule):
+    """AlphaBetaRule with alpha=2, beta=1"""
+    def __init__(self, *args, **kwargs):
+        super(Alpha2Beta1Rule, self).__init__(*args,
+                                              alpha=2,
+                                              beta=1,
+                                              bias=True,
+                                              **kwargs)
+
+
+class Alpha2Beta1IgnoreBiasRule(AlphaBetaRule):
+    """AlphaBetaRule with alpha=2, beta=1 and ignores biases"""
+    def __init__(self, *args, **kwargs):
+        super(Alpha2Beta1IgnoreBiasRule, self).__init__(*args,
+                                                        alpha=2,
+                                                        beta=1,
+                                                        bias=False,
+                                                        **kwargs)
+
+
+class Alpha1Beta0Rule(AlphaBetaRule):
+    """AlphaBetaRule with alpha=1, beta=0"""
+    def __init__(self, *args, **kwargs):
+        super(Alpha1Beta0Rule, self).__init__(*args,
+                                              alpha=1,
+                                              beta=0,
+                                              bias=True,
+                                              **kwargs)
+
+
+class Alpha1Beta0IgnoreBiasRule(AlphaBetaRule):
+    """AlphaBetaRule with alpha=1, beta=0 and ignores biases"""
+    def __init__(self, *args, **kwargs):
+        super(Alpha1Beta0IgnoreBiasRule, self).__init__(*args,
+                                                        alpha=1,
+                                                        beta=0,
+                                                        bias=False,
+                                                        **kwargs)
+
+
+class AlphaBetaXRule(kgraph.ReverseMappingBase):
+    """
+    AlphaBeta advanced as proposed by Alexander Binder.
+    """
+
+    def __init__(self,
+                 layer,
+                 state,
+                 alpha=(0.5, 0.5),
+                 beta=(0.5, 0.5),
+                 bias=True,
+                 copy_weights=False):
+        self._alpha = alpha
+        self._beta = beta
+
+        # prepare positive and negative weights for computing positive
+        # and negative preactivations z in apply_accordingly.
+        if copy_weights:
+            weights = layer.get_weights()
+            if not bias and getattr(layer, "use_bias", False):
+                weights = weights[:-1]
+            positive_weights = [x * (x > 0) for x in weights]
+            negative_weights = [x * (x < 0) for x in weights]
+        else:
+            weights = layer.weights
+            if not bias and getattr(layer, "use_bias", False):
+                weights = weights[:-1]
+            positive_weights = [x * iK.to_floatx(x > 0) for x in weights]
+            negative_weights = [x * iK.to_floatx(x < 0) for x in weights]
+
+        self._layer_wo_act_positive = kgraph.copy_layer_wo_activation(
+            layer,
+            keep_bias=bias,
+            weights=positive_weights,
+            name_template="reversed_kernel_positive_%s")
+        self._layer_wo_act_negative = kgraph.copy_layer_wo_activation(
+            layer,
+            keep_bias=bias,
+            weights=negative_weights,
+            name_template="reversed_kernel_negative_%s")
+
+    def apply(self, Xs, Ys, Rs, reverse_state):
+        #this method is correct, but wasteful
+        grad = ilayers.GradientWRT(len(Xs))
+        times_alpha0 = keras.layers.Lambda(lambda x: x * self._alpha[0])
+        times_alpha1 = keras.layers.Lambda(lambda x: x * self._alpha[1])
+        times_beta0 = keras.layers.Lambda(lambda x: x * self._beta[0])
+        times_beta1 = keras.layers.Lambda(lambda x: x * self._beta[1])
+        keep_positives = keras.layers.Lambda(
+            lambda x: x * K.cast(K.greater(x,0), K.floatx()))
+        keep_negatives = keras.layers.Lambda(
+            lambda x: x * K.cast(K.less(x,0), K.floatx()))
+
+        def f(layer, X):
+            Zs = kutils.apply(layer, X)
+            # Divide incoming relevance by the activations.
+            tmp = [ilayers.SafeDivide()([a, b])
+                    for a, b in zip(Rs, Zs)]
+            # Propagate the relevance to the input neurons
+            # using the gradient
+            tmp = iutils.to_list(grad(X+Zs+tmp))
+            # Re-weight relevance with the input values.
+            tmp = [keras.layers.Multiply()([a, b])
+                    for a, b in zip(X, tmp)]
+            return tmp
+
+        # Distinguish postive and negative inputs.
+        Xs_pos = kutils.apply(keep_positives, Xs)
+        Xs_neg = kutils.apply(keep_negatives, Xs)
+
+        # xpos*wpos
+        r_pp = f(self._layer_wo_act_positive, Xs_pos)
+        # xneg*wneg
+        r_nn = f(self._layer_wo_act_negative, Xs_neg)
+        # a0 * r_pp + a1 * r_nn
+        r_pos = [keras.layers.Add()([times_alpha0(pp), times_beta1(nn)])
+                 for pp, nn in zip(r_pp, r_nn)]
+
+        # xpos*wneg
+        r_pn = f(self._layer_wo_act_negative, Xs_pos)
+        # xneg*wpos
+        r_np = f(self._layer_wo_act_positive, Xs_neg)
+        # b0 * r_pn + b1 * r_np
+        r_neg = [keras.layers.Add()([times_beta0(pn), times_beta1(np)])
+                 for pn, np in zip(r_pn, r_np)]
+
+        return [keras.layers.Subtract()([a, b]) for a, b in zip(r_pos, r_neg)]
+
+
+class AlphaBetaX1000Rule(AlphaBetaXRule):
+    def __init__(self, *args, **kwargs):
+        super(AlphaBetaX1000Rule, self).__init__(*args,
+                                                 alpha=(1, 0),
+                                                 beta=(0, 0),
+                                                 bias=True,
+                                                 **kwargs)
+
+
+class AlphaBetaX1010Rule(AlphaBetaXRule):
+    def __init__(self, *args, **kwargs):
+        super(AlphaBetaX1010Rule, self).__init__(*args,
+                                                 alpha=(1, 0),
+                                                 beta=(0, -1),
+                                                 bias=True,
+                                                 **kwargs)
+
+
+class AlphaBetaX1001Rule(AlphaBetaXRule):
+    def __init__(self, *args, **kwargs):
+        super(AlphaBetaX1001Rule, self).__init__(*args,
+                                                 alpha=(1, 1),
+                                                 beta=(0, 0),
+                                                 bias=True,
+                                                 **kwargs)
+
+
+class AlphaBetaX2m100Rule(AlphaBetaXRule):
+    def __init__(self, *args, **kwargs):
+        super(AlphaBetaX2m100Rule, self).__init__(*args,
+                                                  alpha=(2, 0),
+                                                  beta=(1, 0),
+                                                  bias=True,
+                                                  **kwargs)
+
+
+class BoundedRule(kgraph.ReverseMappingBase):
+    """Z_B rule from the Deep Taylor Decomposition"""
+    # TODO: this only works for relu networks, needs to be extended.
+    # TODO: check
+    def __init__(self, layer, state, low=-1, high=1, copy_weights=False):
+        self._low = low
+        self._high = high
+
+        # This rule works with three variants of the layer, all without biases.
+        # One is the original form and two with only the positive or
+        # negative weights.
+        if copy_weights:
+            weights = layer.get_weights()
+            if layer.use_bias:
+                weights = weights[:-1]
+            positive_weights = [x * (x > 0) for x in weights]
+            negative_weights = [x * (x < 0) for x in weights]
+        else:
+            weights = layer.weights
+            if layer.use_bias:
+                weights = weights[:-1]
+            positive_weights = [x * iK.to_floatx(x > 0) for x in weights]
+            negative_weights = [x * iK.to_floatx(x < 0) for x in weights]
+
+        self._layer_wo_act = kgraph.copy_layer_wo_activation(
+            layer,
+            keep_bias=False,
+            name_template="reversed_kernel_%s")
+        self._layer_wo_act_positive = kgraph.copy_layer_wo_activation(
+            layer,
+            keep_bias=False,
+            weights=positive_weights,
+            name_template="reversed_kernel_positive_%s")
+        self._layer_wo_act_negative = kgraph.copy_layer_wo_activation(
+            layer,
+            keep_bias=False,
+            weights=negative_weights,
+            name_template="reversed_kernel_negative_%s")
+
+    # TODO: clean up this implementation and add more documentation
+    def apply(self, Xs, Ys, Rs, reverse_state):
+        grad = ilayers.GradientWRT(len(Xs))
+        to_low = keras.layers.Lambda(lambda x: x * 0 + self._low)
+        to_high = keras.layers.Lambda(lambda x: x * 0 + self._high)
+
+        low = [to_low(x) for x in Xs]
+        high = [to_high(x) for x in Xs]
+
+        # Get values for the division.
+        A = kutils.apply(self._layer_wo_act, Xs)
+        B = kutils.apply(self._layer_wo_act_positive, low)
+        C = kutils.apply(self._layer_wo_act_negative, high)
+        Zs = [keras.layers.Subtract()([a, keras.layers.Add()([b, c])])
+              for a, b, c in zip(A, B, C)]
+
+        # Divide relevances with the value.
+        tmp = [ilayers.SafeDivide()([a, b])
+               for a, b in zip(Rs, Zs)]
+        # Distribute along the gradient.
+        tmpA = iutils.to_list(grad(Xs+A+tmp))
+        tmpB = iutils.to_list(grad(low+B+tmp))
+        tmpC = iutils.to_list(grad(high+C+tmp))
+
+        tmpA = [keras.layers.Multiply()([a, b]) for a, b in zip(Xs, tmpA)]
+        tmpB = [keras.layers.Multiply()([a, b]) for a, b in zip(low, tmpB)]
+        tmpC = [keras.layers.Multiply()([a, b]) for a, b in zip(high, tmpC)]
+
+        tmp = [keras.layers.Subtract()([a, keras.layers.Add()([b, c])])
+               for a, b, c in zip(tmpA, tmpB, tmpC)]
+
+        return tmp
+
+
+
+class ZPlusRule(Alpha1Beta0IgnoreBiasRule):
+    """
+    The ZPlus rule is a special case of the AlphaBetaRule
+    for alpha=1, beta=0, which assumes inputs x >= 0
+    and ignores the bias.
+    CAUTION! Results differ from Alpha=1, Beta=0
+    if inputs are not strictly >= 0
+    """
+    #TODO: assert that layer inputs are always >= 0
+    def __init__(self, *args, **kwargs):
+        super(ZPlusRule, self).__init__(*args, **kwargs)
+
+
+
+class ZPlusFastRule(kgraph.ReverseMappingBase):
+    """
+    The ZPlus rule is a special case of the AlphaBetaRule
+    for alpha=1, beta=0 and assumes inputs x >= 0.
+    """
+
+    def __init__(self, layer, state, copy_weights=False):
+        # The z-plus rule only works with positive weights and
+        # no biases.
+        #TODO: assert that layer inputs are always >= 0
+        if copy_weights:
+            weights = layer.get_weights()
+            if layer.use_bias:
+                weights = weights[:-1]
+            weights = [x * (x > 0) for x in weights]
+        else:
+            weights = layer.weights
+            if layer.use_bias:
+                weights = weights[:-1]
+            weights = [x * iK.to_floatx(x > 0) for x in weights]
+
+        self._layer_wo_act_b_positive = kgraph.copy_layer_wo_activation(
+            layer,
+            keep_bias=False,
+            weights=weights,
+            name_template="reversed_kernel_positive_%s")
+
+    def apply(self, Xs, Ys, Rs, reverse_state):
+        grad = ilayers.GradientWRT(len(Xs))
+
+        #TODO: assert all inputs are positive, instead of only keeping the positives.
+        #keep_positives = keras.layers.Lambda(lambda x: x * K.cast(K.greater(x,0), K.floatx()))
+        #Xs = kutils.apply(keep_positives, Xs)
+
+        # Get activations.
+        Zs = kutils.apply(self._layer_wo_act_b_positive, Xs)
+        # Divide incoming relevance by the activations.
+        tmp = [ilayers.SafeDivide()([a, b])
+               for a, b in zip(Rs, Zs)]
+        # Propagate the relevance to input neurons
+        # using the gradient.
+        tmp = iutils.to_list(grad(Xs+Zs+tmp))
+        # Re-weight relevance with the input values.
+        return [keras.layers.Multiply()([a, b])
+                for a, b in zip(Xs, tmp)]

+ 102 - 0
original_model/innvestigate/analyzer/relevance_based/utils.py

@@ -0,0 +1,102 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+__all__ = [
+    "assert_lrp_epsilon_param",
+    "assert_infer_lrp_alpha_beta_param"
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+def assert_lrp_epsilon_param(epsilon, caller):
+    """
+        Function for asserting epsilon parameter choice
+        passed to constructors inheriting from EpsilonRule
+        and LRPEpsilon.
+        The following conditions can not be met:
+
+        epsilon > 1
+
+        :param epsilon: the epsilon parameter.
+        :param caller: the class instance calling this assertion function
+    """
+
+    if epsilon <= 0:
+        err_head = "Constructor call to {} : ".format(caller.__class__.__name__)
+        err_msg = err_head + "Parameter epsilon must be > 0 but was {}".format(epsilon)
+        raise ValueError(err_msg)
+    return epsilon
+
+
+def assert_infer_lrp_alpha_beta_param(alpha, beta, caller):
+    """
+        Function for asserting parameter choices for alpha and beta
+        passed to constructors inheriting from AlphaBetaRule
+        and LRPAlphaBeta.
+
+        since alpha - beta are subjected to sum to 1,
+        it is sufficient for only one of the parameters to be passed
+        to a corresponding class constructor.
+        this method will cause an assertion error if both are None
+        or the following conditions can not be met
+
+        alpha >= 1
+        beta >= 0
+        alpha - beta = 1
+
+        :param alpha: the alpha parameter.
+        :param beta: the beta parameter
+        :param caller: the class instance calling this assertion function
+    """
+
+    err_head = "Constructor call to {} : ".format(caller.__class__.__name__)
+    if alpha is None and beta is None:
+        err_msg = err_head + "Neither alpha or beta were given"
+        raise ValueError(err_msg)
+
+    #assert passed parameter choices
+    if alpha is not None and alpha < 1:
+        err_msg = err_head +"Passed parameter alpha invalid. Expecting alpha >= 1 but was {}".format(alpha)
+        raise ValueError(err_msg)
+
+    if beta is not None and beta < 0:
+        err_msg = err_head +"Passed parameter beta invalid. Expecting beta >= 0 but was {}".format(beta)
+        raise ValueError(err_msg)
+
+    #assert inferred parameter choices
+    if alpha is None:
+        alpha = beta + 1
+        if alpha < 1:
+            err_msg = err_head +"Inferring alpha from given beta {} s.t. alpha - beta = 1, with condition alpha >= 1 not possible.".format(beta)
+            raise ValueError(err_msg)
+
+
+    if beta is None:
+        beta = alpha - 1
+        if beta < 0:
+            err_msg = err_head +"Inferring beta from given alpha {} s.t. alpha - beta = 1, with condition beta >= 0 not possible.".format(alpha)
+            raise ValueError(err_msg)
+
+
+    #final check: alpha - beta = 1
+    amb = alpha - beta
+    if amb != 1:
+        err_msg = err_head +"Condition alpha - beta = 1 not fulfilled. alpha={} ; beta={} -> alpha - beta = {}".format(alpha, beta, amb)
+        raise ValueError(err_msg)
+
+    #return benign values for alpha and beta
+    return alpha, beta
+
+###############################################################################
+###############################################################################
+###############################################################################

+ 333 - 0
original_model/innvestigate/analyzer/wrapper.py

@@ -0,0 +1,333 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import zip
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.models
+import keras.backend as K
+import numpy as np
+
+
+from . import base
+from .. import layers as ilayers
+from .. import utils as iutils
+from ..utils import keras as kutils
+
+
+__all__ = [
+    "WrapperBase",
+    "AugmentReduceBase",
+    "GaussianSmoother",
+    "PathIntegrator",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class WrapperBase(base.AnalyzerBase):
+    """Interface for wrappers around analyzers
+
+    This class is the basic interface for wrappers around analyzers.
+
+    :param subanalyzer: The analyzer to be wrapped.
+    """
+
+    def __init__(self, subanalyzer, *args, **kwargs):
+        self._subanalyzer = subanalyzer
+        model = None
+
+        super(WrapperBase, self).__init__(model,
+                                          *args, **kwargs)
+
+    def analyze(self, *args, **kwargs):
+        return self._subanalyzer.analyze(*args, **kwargs)
+
+    def _get_state(self):
+        sa_class_name, sa_state = self._subanalyzer.save()
+
+        state = {}
+        state.update({"subanalyzer_class_name": sa_class_name})
+        state.update({"subanalyzer_state": sa_state})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        sa_class_name = state.pop("subanalyzer_class_name")
+        sa_state = state.pop("subanalyzer_state")
+        assert len(state) == 0
+
+        subanalyzer = base.AnalyzerBase.load(sa_class_name, sa_state)
+        kwargs = {"subanalyzer": subanalyzer}
+        return kwargs
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class AugmentReduceBase(WrapperBase):
+    """Interface for wrappers that augment the input and reduce the analysis.
+
+    This class is an interface for wrappers that:
+    * augment the input to the analyzer by creating new samples.
+    * reduce the returned analysis to match the initial input shapes.
+
+    :param subanalyzer: The analyzer to be wrapped.
+    :param augment_by_n: Number of samples to create.
+    """
+
+    def __init__(self, subanalyzer, *args, **kwargs):
+        self._augment_by_n = kwargs.pop("augment_by_n", 2)
+        self._neuron_selection_mode = subanalyzer._neuron_selection_mode
+
+        if self._neuron_selection_mode != "all":
+            # TODO: this is not transparent, find a better way.
+            subanalyzer._neuron_selection_mode = "index"
+        super(AugmentReduceBase, self).__init__(subanalyzer,
+                                                *args, **kwargs)
+
+        if isinstance(self._subanalyzer, base.AnalyzerNetworkBase):
+            # Take the keras analyzer model and
+            # add augment and reduce functionality.
+            self._keras_based_augment_reduce = True
+        else:
+            raise NotImplementedError("Keras-based subanalyzer required.")
+
+    def create_analyzer_model(self):
+        if not self._keras_based_augment_reduce:
+            return
+
+        self._subanalyzer.create_analyzer_model()
+
+        if self._subanalyzer._n_debug_output > 0:
+            raise Exception("No debug output at subanalyzer is supported.")
+
+        model = self._subanalyzer._analyzer_model
+        if None in model.input_shape[1:]:
+            raise ValueError("The input shape for the model needs "
+                             "to be fully specified (except the batch axis). "
+                             "Model input shape is: %s" % (model.input_shape,))
+
+        inputs = model.inputs[:self._subanalyzer._n_data_input]
+        extra_inputs = model.inputs[self._subanalyzer._n_data_input:]
+        # todo: check this, index seems not right.
+        #outputs = model.outputs[:self._subanalyzer._n_data_input]
+        extra_outputs = model.outputs[self._subanalyzer._n_data_input:]
+
+        if len(extra_outputs) > 0:
+            raise Exception("No extra output is allowed "
+                            "with this wrapper.")
+
+        new_inputs = iutils.to_list(self._augment(inputs))
+        # print(type(new_inputs), type(extra_inputs))
+        tmp = iutils.to_list(model(new_inputs+extra_inputs))
+        new_outputs = iutils.to_list(self._reduce(tmp))
+        new_constant_inputs = self._keras_get_constant_inputs()
+
+        new_model = keras.models.Model(
+            inputs=inputs+extra_inputs+new_constant_inputs,
+            outputs=new_outputs+extra_outputs)
+        self._subanalyzer._analyzer_model = new_model
+
+    def analyze(self, X, *args, **kwargs):
+        if self._keras_based_augment_reduce is True:
+            if not hasattr(self._subanalyzer, "_analyzer_model"):
+                self.create_analyzer_model()
+
+            ns_mode = self._neuron_selection_mode
+            if ns_mode in ["max_activation", "index"]:
+                if ns_mode == "max_activation":
+                    tmp = self._subanalyzer._model.predict(X)
+                    indices = np.argmax(tmp, axis=1)
+                else:
+                    if len(args):
+                        args = list(args)
+                        indices = args.pop(0)
+                    else:
+                        indices = kwargs.pop("neuron_selection")
+
+                # broadcast to match augmented samples.
+                indices = np.repeat(indices, self._augment_by_n)
+
+                kwargs["neuron_selection"] = indices
+            return self._subanalyzer.analyze(X, *args, **kwargs)
+        else:
+            raise DeprecationWarning("Not supported anymore.")
+
+    def _keras_get_constant_inputs(self):
+        return list()
+
+    def _augment(self, X):
+        repeat = ilayers.Repeat(self._augment_by_n, axis=0)
+        return [repeat(x) for x in iutils.to_list(X)]
+
+    def _reduce(self, X):
+        X_shape = [K.int_shape(x) for x in iutils.to_list(X)]
+        reshape = [ilayers.Reshape((-1, self._augment_by_n)+shape[1:])
+                   for shape in X_shape]
+        mean = ilayers.Mean(axis=1)
+
+        return [mean(reshape_x(x)) for x, reshape_x in zip(X, reshape)]
+
+    def _get_state(self):
+        if self._neuron_selection_mode != "all":
+            # TODO: this is not transparent, find a better way.
+            # revert the tempering in __init__
+            tmp = self._neuron_selection_mode
+            self._subanalyzer._neuron_selection_mode = tmp
+        state = super(AugmentReduceBase, self)._get_state()
+        state.update({"augment_by_n": self._augment_by_n})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        augment_by_n = state.pop("augment_by_n")
+        kwargs = super(AugmentReduceBase, clazz)._state_to_kwargs(state)
+        kwargs.update({"augment_by_n": augment_by_n})
+        return kwargs
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class GaussianSmoother(AugmentReduceBase):
+    """Wrapper that adds noise to the input and averages over analyses
+
+    This wrapper creates new samples by adding Gaussian noise
+    to the input. The final analysis is an average of the returned analyses.
+
+    :param subanalyzer: The analyzer to be wrapped.
+    :param noise_scale: The stddev of the applied noise.
+    :param augment_by_n: Number of samples to create.
+    """
+
+    def __init__(self, subanalyzer, *args, **kwargs):
+        self._noise_scale = kwargs.pop("noise_scale", 1)
+        super(GaussianSmoother, self).__init__(subanalyzer,
+                                               *args, **kwargs)
+
+    def _augment(self, X):
+        tmp = super(GaussianSmoother, self)._augment(X)
+        noise = ilayers.TestPhaseGaussianNoise(stddev=self._noise_scale)
+        return [noise(x) for x in tmp]
+
+    def _get_state(self):
+        state = super(GaussianSmoother, self)._get_state()
+        state.update({"noise_scale": self._noise_scale})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        noise_scale = state.pop("noise_scale")
+        kwargs = super(GaussianSmoother, clazz)._state_to_kwargs(state)
+        kwargs.update({"noise_scale": noise_scale})
+        return kwargs
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class PathIntegrator(AugmentReduceBase):
+    """Integrated the analysis along a path
+
+    This analyzer:
+    * creates a path from input to reference image.
+    * creates steps number of intermediate inputs and
+      crests an analysis for them.
+    * sums the analyses and multiplies them with the input-reference_input.
+
+    This wrapper is used to implement Integrated Gradients.
+    We refer to the paper for further information.
+
+    :param subanalyzer: The analyzer to be wrapped.
+    :param steps: Number of steps for integration.
+    :param reference_inputs: The reference input.
+    """
+
+    def __init__(self, subanalyzer, *args, **kwargs):
+        steps = kwargs.pop("steps", 16)
+        self._reference_inputs = kwargs.pop("reference_inputs", 0)
+        self._keras_constant_inputs = None
+        super(PathIntegrator, self).__init__(subanalyzer,
+                                             *args,
+                                             augment_by_n=steps,
+                                             **kwargs)
+
+    def _keras_set_constant_inputs(self, inputs):
+        tmp = [K.variable(x) for x in inputs]
+        self._keras_constant_inputs = [
+            keras.layers.Input(tensor=x, shape=x.shape[1:])
+            for x in tmp]
+
+    def _keras_get_constant_inputs(self):
+        return self._keras_constant_inputs
+
+    def _compute_difference(self, X):
+        if self._keras_constant_inputs is None:
+            tmp = kutils.broadcast_np_tensors_to_keras_tensors(
+                X, self._reference_inputs)
+            self._keras_set_constant_inputs(tmp)
+
+        reference_inputs = self._keras_get_constant_inputs()
+        return [keras.layers.Subtract()([x, ri])
+                for x, ri in zip(X, reference_inputs)]
+
+    def _augment(self, X):
+        tmp = super(PathIntegrator, self)._augment(X)
+        tmp = [ilayers.Reshape((-1, self._augment_by_n)+K.int_shape(x)[1:])(x)
+               for x in tmp]
+
+        difference = self._compute_difference(X)
+        self._keras_difference = difference
+        # Make broadcastable.
+        difference = [ilayers.Reshape((-1, 1)+K.int_shape(x)[1:])(x)
+                      for x in difference]
+
+        # Compute path steps.
+        multiply_with_linspace = ilayers.MultiplyWithLinspace(
+            0, 1,
+            n=self._augment_by_n,
+            axis=1)
+        path_steps = [multiply_with_linspace(d) for d in difference]
+
+        reference_inputs = self._keras_get_constant_inputs()
+        ret = [keras.layers.Add()([x, p]) for x, p in zip(reference_inputs, path_steps)]
+        ret = [ilayers.Reshape((-1,)+K.int_shape(x)[2:])(x) for x in ret]
+        return ret
+
+    def _reduce(self, X):
+        tmp = super(PathIntegrator, self)._reduce(X)
+        difference = self._keras_difference
+        del self._keras_difference
+
+        return [keras.layers.Multiply()([x, d])
+                for x, d in zip(tmp, difference)]
+
+    def _get_state(self):
+        state = super(PathIntegrator, self)._get_state()
+        state.update({"reference_inputs": self._reference_inputs})
+        return state
+
+    @classmethod
+    def _state_to_kwargs(clazz, state):
+        reference_inputs = state.pop("reference_inputs")
+        kwargs = super(PathIntegrator, clazz)._state_to_kwargs(state)
+        kwargs.update({"reference_inputs": reference_inputs})
+        # We use steps instead.
+        kwargs.update({"steps": kwargs["augment_by_n"]})
+        del kwargs["augment_by_n"]
+        return kwargs

+ 0 - 0
original_model/innvestigate/applications/__init__.py


+ 297 - 0
original_model/innvestigate/applications/imagenet.py

@@ -0,0 +1,297 @@
+"""Example applications for image classifcation.
+
+Each function returns a pretrained ImageNet model.
+The models are based on keras.applications models and
+contain additionally pretrained patterns.
+
+The returned dictionary contains the following
+keys\: model, in, sm_out, out, image_shape, color_coding,
+preprocess_f, patterns.
+
+Function parameters\:
+
+:param load_weights: Download or access cached weights.
+:param load_patterns: Download or access cached patterns.
+"""
+# todo: rename in, sm_out, out to input_tensors, output_tensors,
+# todo: softmax_output_tenors
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import range
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.backend as K
+import keras.applications.resnet50
+import keras.applications.vgg16
+import keras.applications.vgg19
+import keras.applications.inception_v3
+import keras.applications.inception_resnet_v2
+import keras.applications.densenet
+import keras.applications.nasnet
+import keras.utils.data_utils
+import numpy as np
+import warnings
+
+from ..utils.keras import graph as kgraph
+
+
+__all__ = [
+    "vgg16",
+    "vgg19",
+    "resnet50",
+    "inception_v3",
+    "inception_resnet_v2",
+    "densenet121",
+    "densenet169",
+    "densenet201",
+    "nasnet_large",
+    "nasnet_mobile",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+PATTERNS = {
+    "vgg16_pattern_type_relu_tf_dim_ordering_tf_kernels.npz": {
+        "url": "https://www.dropbox.com/s/15lip81fzvbgkaa/vgg16_pattern_type_relu_tf_dim_ordering_tf_kernels.npz?dl=1",
+        "hash": "8c2abe648e116a93fd5027fab49177b0",
+    },
+    "vgg19_pattern_type_relu_tf_dim_ordering_tf_kernels.npz": {
+        "url": "https://www.dropbox.com/s/nc5empj78rfe9hm/vgg19_pattern_type_relu_tf_dim_ordering_tf_kernels.npz?dl=1",
+        "hash": "3258b6c64537156afe75ca7b3be44742",
+    },
+}
+
+
+def _get_patterns_info(netname, pattern_type):
+    if pattern_type is True:
+        pattern_type = "relu"
+
+    file_name = ("%s_pattern_type_%s_tf_dim_ordering_tf_kernels.npz" %
+                 (netname, pattern_type))
+
+    return {"file_name": file_name,
+            "url": PATTERNS[file_name]["url"],
+            "hash": PATTERNS[file_name]["hash"]}
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def _prepare_keras_net(netname,
+                       clazz,
+                       image_shape,
+                       preprocess_f,
+                       preprocess_mode=None,
+                       color_coding="RGB",
+                       load_weights=False,
+                       load_patterns=False):
+    net = {}
+    net["name"] = netname
+    net["image_shape"] = image_shape
+    if K.image_data_format() == "channels_first":
+        net["input_shape"] = [None, 3]+image_shape
+    else:
+        net["input_shape"] = [None]+image_shape+[3]
+
+    weights = None
+    if load_weights is True:
+        weights = "imagenet"
+
+    model = clazz(weights=weights,
+                  input_shape=tuple(net["input_shape"][1:]))
+    net["model"] = model
+
+    net["in"] = model.inputs
+    net["sm_out"] = model.outputs
+    net["out"] = kgraph.pre_softmax_tensors(model.outputs)
+
+    net["color_coding"] = color_coding
+    net["preprocess_f"] = preprocess_f
+    net["input_range"] = {
+        None: (-128, 128),
+        "caffe": (-128, 128),
+        "tf": (-1, 1),
+        "torch": (-3, 3),
+    }[preprocess_mode]
+
+    net["patterns"] = None
+    if load_patterns is not False:
+        try:
+            pattern_info = _get_patterns_info(netname, load_patterns)
+        except KeyError:
+            warnings.warn("There are no patterns for network '%s'." % netname)
+        else:
+            patterns_path = keras.utils.data_utils.get_file(
+                pattern_info["file_name"],
+                pattern_info["url"],
+                cache_subdir="innvestigate_patterns",
+                hash_algorithm="md5",
+                file_hash=pattern_info["hash"])
+            patterns_file = np.load(patterns_path)
+            patterns = [patterns_file["arr_%i" % i]
+                        for i in range(len(patterns_file.keys()))]
+            net["patterns"] = patterns
+    return net
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def vgg16(load_weights=False, load_patterns=False):
+    return _prepare_keras_net(
+        "vgg16",
+        keras.applications.vgg16.VGG16,
+        [224, 224],
+        preprocess_f=keras.applications.vgg16.preprocess_input,
+        preprocess_mode="caffe",
+        color_coding="BGR",
+        load_weights=load_weights,
+        load_patterns=load_patterns)
+
+
+def vgg19(load_weights=False, load_patterns=False):
+    return _prepare_keras_net(
+        "vgg19",
+        keras.applications.vgg19.VGG19,
+        [224, 224],
+        preprocess_f=keras.applications.vgg19.preprocess_input,
+        preprocess_mode="caffe",
+        color_coding="BGR",
+        load_weights=load_weights,
+        load_patterns=load_patterns)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def resnet50(load_weights=False, load_patterns=False):
+    return _prepare_keras_net(
+        "resnet50",
+        keras.applications.resnet50.ResNet50,
+        [224, 224],
+        preprocess_f=keras.applications.resnet50.preprocess_input,
+        preprocess_mode="caffe",
+        color_coding="BGR",
+        load_weights=load_weights,
+        load_patterns=load_patterns)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def inception_v3(load_weights=False, load_patterns=False):
+    return _prepare_keras_net(
+        "inception_v3",
+        keras.applications.inception_v3.InceptionV3,
+        [299, 299],
+        preprocess_f=keras.applications.inception_v3.preprocess_input,
+        preprocess_mode="tf",
+        load_weights=load_weights,
+        load_patterns=load_patterns)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def inception_resnet_v2(load_weights=False, load_patterns=False):
+    return _prepare_keras_net(
+        "inception_resnet_v2",
+        keras.applications.inception_resnet_v2.InceptionResNetV2,
+        [299, 299],
+        preprocess_f=keras.applications.inception_resnet_v2.preprocess_input,
+        preprocess_mode="tf",
+        load_weights=load_weights,
+        load_patterns=load_patterns)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def densenet121(load_weights=False, load_patterns=False):
+    return _prepare_keras_net(
+        "densenet121",
+        keras.applications.densenet.DenseNet121,
+        [224, 224],
+        preprocess_f=keras.applications.densenet.preprocess_input,
+        preprocess_mode="torch",
+        load_weights=load_weights,
+        load_patterns=load_patterns)
+
+
+def densenet169(load_weights=False, load_patterns=False):
+    return _prepare_keras_net(
+        "densenet169",
+        keras.applications.densenet.DenseNet169,
+        [224, 224],
+        preprocess_f=keras.applications.densenet.preprocess_input,
+        preprocess_mode="torch",
+        load_weights=load_weights,
+        load_patterns=load_patterns)
+
+
+def densenet201(load_weights=False, load_patterns=False):
+    return _prepare_keras_net(
+        "densenet201",
+        keras.applications.densenet.DenseNet201,
+        [224, 224],
+        preprocess_f=keras.applications.densenet.preprocess_input,
+        preprocess_mode="torch",
+        load_weights=load_weights,
+        load_patterns=load_patterns)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def nasnet_large(load_weights=False, load_patterns=False):
+    if K.image_data_format() == "channels_first":
+        raise Exception("NASNet is not available for channels first.")
+
+    return _prepare_keras_net(
+        "nasnet_large",
+        keras.applications.nasnet.NASNetLarge,
+        [331, 331],
+        color_coding="BGR",
+        preprocess_f=keras.applications.nasnet.preprocess_input,
+        preprocess_mode="tf",
+        load_weights=load_weights,
+        load_patterns=load_patterns)
+
+
+def nasnet_mobile(load_weights=False, load_patterns=False):
+    if K.image_data_format() == "channels_first":
+        raise Exception("NASNet is not available for channels first.")
+
+    return _prepare_keras_net(
+        "nasnet_mobile",
+        keras.applications.nasnet.NASNetMobile,
+        [224, 224],
+        color_coding="BGR",
+        preprocess_f=keras.applications.nasnet.preprocess_input,
+        preprocess_mode="tf",
+        load_weights=load_weights,
+        load_patterns=load_patterns)

+ 106 - 0
original_model/innvestigate/applications/mnist.py

@@ -0,0 +1,106 @@
+"""Example applications for image classifcation.
+
+Each function returns a pretrained MNIST model.
+The models are based on https://doi.org/10.1371/journal.pone.0130140
+and http://jmlr.org/papers/v17/15-618.html.
+
+"""
+# TODO: rename in, sm_out, out to input_tensors, output_tensors,
+# TODO: softmax_output_tenors
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import os
+import keras.utils.data_utils
+import numpy as np
+
+import keras.models
+from keras.models import load_model, clone_model
+#from keras.utils import get_file
+
+
+
+__all__ = [
+    "pretrained_plos_long_relu",
+    "pretrained_plos_short_relu",
+    "pretrained_plos_long_tanh",
+    "pretrained_plos_short_tanh",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+# pre-trained models from [https://doi.org/10.1371/journal.pone.0130140 , http://jmlr.org/papers/v17/15-618.html]
+PRETRAINED_MODELS = {"pretrained_plos_long_relu":
+                        {"file":"plos-mnist-rect-long.h5",
+                         "url" : "https://www.dropbox.com/s/26w7i58qqcuosn4/plos-mnist-rect-long.h5"
+                        },
+                     "pretrained_plos_short_relu":
+                        {"file":"plos-mnist-rect-short.h5",
+                         "url":"https://www.dropbox.com/s/89nvwyls55xycmw/plos-mnist-rect-short.h5"
+                        },
+                     "pretrained_plos_long_tanh":
+                        {"file":"plos-mnist-tanh-long.h5",
+                         "url":"https://www.dropbox.com/s/61e3a4gdbjo9bca/plos-mnist-tanh-long.h5"
+                        },
+                     "pretrained_plos_short_tanh":
+                        {"file":"plos-mnist-tanh-short.h5",
+                         "url":"https://www.dropbox.com/s/foqv60kot0retfr/plos-mnist-tanh-short.h5"
+                        },
+                    }
+
+
+def _load_pretrained_net(modelname, new_input_shape):
+    filename = PRETRAINED_MODELS[modelname]["file"]
+    urlname = PRETRAINED_MODELS[modelname]["url"]
+    #model_path = get_file(fname=filename, origin=urlname) #TODO: FIX! corrupts the file?
+    model_path = os.path.expanduser('~') + "/.keras/models/" + filename
+
+
+    #workaround the more elegant, but dysfunctional solution.
+    if not os.path.isfile(model_path):
+        model_dir = os.path.dirname(model_path)
+        if not os.path.isdir(model_dir):
+            os.makedirs(model_dir)
+        os.system("wget {} &&  mv -v {} {}".format(urlname, filename, model_path))
+
+
+    model = load_model(model_path)
+    #create replacement input layer with new shape.
+    model.layers[0] = keras.layers.InputLayer(input_shape=new_input_shape, name="input_1")
+    for l in model.layers:
+        l.name = "%s_workaround" % l.name
+    model = keras.models.Sequential(layers=model.layers)
+
+    model_w_sm = clone_model(model)
+
+    #NOTE: perform forward pass to fix a keras 2.2.0 related issue with improper weight initialization
+    #See: https://github.com/albermax/innvestigate/issues/88
+    x_dummy = np.zeros(new_input_shape)[None, ...]
+    model_w_sm.predict(x_dummy)
+
+    model_w_sm.set_weights(model.get_weights())
+    model_w_sm.add(keras.layers.Activation("softmax"))
+    return model, model_w_sm
+
+
+def pretrained_plos_long_relu(input_shape, **kwargs):
+    return _load_pretrained_net("pretrained_plos_long_relu", input_shape)
+
+def pretrained_plos_short_relu(input_shape, **kwargs):
+    return _load_pretrained_net("pretrained_plos_short_relu", input_shape)
+
+def pretrained_plos_long_tanh(input_shape, **kwargs):
+    return _load_pretrained_net("pretrained_plos_long_tanh", input_shape)
+
+def pretrained_plos_short_tanh(input_shape, **kwargs):
+    return _load_pretrained_net("pretrained_plos_short_tanh", input_shape)

+ 660 - 0
original_model/innvestigate/layers.py

@@ -0,0 +1,660 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import range, zip
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+import keras
+import keras.backend as K
+import keras.constraints
+import keras.layers
+import keras.regularizers
+from keras.utils import conv_utils
+import numpy as np
+
+
+from . import utils as iutils
+from .utils.keras import backend as iK
+
+
+__all__ = [
+    "Constant",
+    "Zero",
+    "One",
+    "ZerosLike",
+    "OnesLike",
+    "AsFloatX",
+    "FiniteCheck",
+
+    "Gradient",
+    "GradientWRT",
+
+    "Min",
+    "Max",
+    "Greater",
+    "Less",
+    "GreaterThanZero",
+    "LessThanZero",
+    "GreaterEqual",
+    "LessEqual",
+    "GreaterEqualThanZero",
+    "LessEqualThanZero",
+    "Sum",
+    "Mean",
+    "CountNonZero",
+
+    "Identity",
+    "Abs",
+    "Square",
+    "Clip",
+    "Project",
+    "Print",
+
+    "Transpose",
+    "Dot",
+    "SafeDivide",
+
+    "Repeat",
+    "Reshape",
+    "MultiplyWithLinspace",
+    "TestPhaseGaussianNoise",
+    "ExtractConv2DPatches",
+    "RunningMeans",
+    "Broadcast",
+    "Gather",
+    "GatherND",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def Constant(c, reference=None):
+    if reference is None:
+        return K.constant(c)
+    else:
+        dtype = K.dtype(reference)
+        return K.constant(np.dtype(dtype)(c), dtype=dtype)
+
+
+def Zero(reference=None):
+    return Constant(0, reference=reference)
+
+
+def One(reference=None):
+    return Constant(1, reference=reference)
+
+
+class ZerosLike(keras.layers.Layer):
+    def call(self, x):
+        return [K.zeros_like(tmp) for tmp in iutils.to_list(x)]
+
+
+class OnesLike(keras.layers.Layer):
+    def call(self, x):
+        return [K.ones_like(tmp) for tmp in iutils.to_list(x)]
+
+
+class AsFloatX(keras.layers.Layer):
+    def call(self, x):
+        return [iK.to_floatx(tmp) for tmp in iutils.to_list(x)]
+
+
+class FiniteCheck(keras.layers.Layer):
+    def call(self, x):
+        return [K.sum(iK.to_floatx(iK.is_not_finite(tmp)))
+                for tmp in iutils.to_list(x)]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class Gradient(keras.layers.Layer):
+    "Returns gradient of sum(output), expects inputs+[output,]."
+
+    def call(self, x):
+        inputs, output = x[:-1], x[-1]
+        return K.gradients(K.sum(output), inputs)
+
+    def compute_output_shape(self, input_shapes):
+        return input_shapes[:-1]
+
+
+class GradientWRT(keras.layers.Layer):
+    "Returns gradient wrt to another layer and given gradient,"
+    " expects inputs+[output,]."
+
+    def __init__(self, n_inputs, mask=None, **kwargs):
+        self.n_inputs = n_inputs
+        self.mask = mask
+        super(GradientWRT, self).__init__(**kwargs)
+
+    def call(self, x):
+        assert isinstance(x, (list, tuple))
+        Xs, tmp_Ys = x[:self.n_inputs], x[self.n_inputs:]
+        assert len(tmp_Ys) % 2 == 0
+        len_Ys = len(tmp_Ys) // 2
+        Ys, known_Ys = tmp_Ys[:len_Ys], tmp_Ys[len_Ys:]
+        ret = iK.gradients(Xs, Ys, known_Ys)
+        if self.mask is not None:
+            ret = [x for c, x in zip(self.mask, ret) if c]
+        self.__workaround__len_ret = len(ret)
+        return ret
+
+    def compute_output_shape(self, input_shapes):
+        if self.mask is None:
+            return input_shapes[:self.n_inputs]
+        else:
+            return [x for c, x in zip(self.mask, input_shapes[:self.n_inputs])
+                    if c]
+
+    # todo: remove once keras is fixed.
+    # this is a workaround for cases when
+    # wrapper and skip connections are used together.
+    # bring the fix into keras and remove once
+    # keras is patched.
+    def compute_mask(self, inputs, mask=None):
+        """Computes an output mask tensor.
+
+        # Arguments
+            inputs: Tensor or list of tensors.
+            mask: Tensor or list of tensors.
+
+        # Returns
+            None or a tensor (or list of tensors,
+                one per output tensor of the layer).
+        """
+        if not self.supports_masking:
+            if mask is not None:
+                if isinstance(mask, list):
+                    if any(m is not None for m in mask):
+                        raise TypeError('Layer ' + self.name +
+                                        ' does not support masking, '
+                                        'but was passed an input_mask: ' +
+                                        str(mask))
+                else:
+                    raise TypeError('Layer ' + self.name +
+                                    ' does not support masking, '
+                                    'but was passed an input_mask: ' +
+                                    str(mask))
+            # masking not explicitly supported: return None as mask
+
+            # this is the workaround for model.run_internal_graph.
+            # it is required that there as many masks as outputs:
+            return [None for _ in range(self.__workaround__len_ret)]
+        # if masking is explicitly supported, by default
+        # carry over the input mask
+        return mask
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class _Reduce(keras.layers.Layer):
+
+    def __init__(self, axis=-1, keepdims=False, *args, **kwargs):
+        self.axis = axis
+        self.keepdims = keepdims
+        super(_Reduce, self).__init__(*args, **kwargs)
+
+    def call(self, x):
+        return self._apply_reduce(x, axis=self.axis, keepdims=self.keepdims)
+
+    def compute_output_shape(self, input_shape):
+        if self.axis is None:
+            if self.keepdims is False:
+                return (1,)
+            else:
+                return tuple(np.ones_like(input_shape))
+        else:
+            axes = np.arange(len(input_shape))
+            if self.keepdims is False:
+                for i in iutils.to_list(self.axis):
+                    axes = np.delete(axes, i, 0)
+            else:
+                for i in iutils.to_list(self.axis):
+                    axes[i] = 1
+            return tuple([idx
+                          for i, idx in enumerate(input_shape)
+                          if i in axes])
+
+    def _apply_reduce(self, x, axis, keepdims):
+        raise NotImplementedError()
+
+
+class Min(_Reduce):
+    def _apply_reduce(self, x, axis, keepdims):
+        return K.min(x, axis=axis, keepdims=keepdims)
+
+
+class Max(_Reduce):
+    def _apply_reduce(self, x, axis, keepdims):
+        return K.max(x, axis=axis, keepdims=keepdims)
+
+
+class Sum(_Reduce):
+    def _apply_reduce(self, x, axis, keepdims):
+        return K.sum(x, axis=axis, keepdims=keepdims)
+
+
+class Mean(_Reduce):
+    def _apply_reduce(self, x, axis, keepdims):
+        return K.mean(x, axis=axis, keepdims=keepdims)
+
+
+class CountNonZero(_Reduce):
+    def _apply_reduce(self, x, axis, keepdims):
+        return K.sum(iK.to_floatx(K.not_equal(x, K.constant(0))),
+                     axis=axis,
+                     keepdims=keepdims)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class _Map(keras.layers.Layer):
+
+    def call(self, x):
+        if isinstance(x, list) and len(x) == 1:
+            x = x[0]
+        return self._apply_map(x)
+
+    def compute_output_shape(self, input_shape):
+        return input_shape
+
+    def _apply_map(self, x):
+        raise NotImplementedError()
+
+
+class Identity(_Map):
+    def _apply_map(self, x):
+        return K.identity(x)
+
+
+class Abs(_Map):
+    def _apply_map(self, x):
+        return K.abs(x)
+
+
+class Square(_Map):
+    def _apply_map(self, x):
+        return K.square(x)
+
+
+class Clip(_Map):
+
+    def __init__(self, min_value, max_value):
+        self._min_value = min_value
+        self._max_value = max_value
+        return super(Clip, self).__init__()
+
+    def _apply_map(self, x):
+        return K.clip(x, self._min_value, self._max_value)
+
+
+class Project(_Map):
+
+    def __init__(self, output_range=False, input_is_postive_only=False):
+        self._output_range = output_range
+        self._input_is_positive_only = input_is_postive_only
+        return super(Project, self).__init__()
+
+    def _apply_map(self, x):
+        def safe_divide(a, b):
+            return a / (b + iK.to_floatx(K.equal(b, K.constant(0))) * 1)
+
+        dims = K.int_shape(x)
+        n_dim = len(dims)
+        axes = tuple(range(1, n_dim))
+        if len(axes) == 1:
+            # TODO(albermax): this is only the case when the dimension in this
+            # axis is 1, fix this.
+            # Cannot reduce
+            return x
+
+        absmax = K.max(K.abs(x),
+                       axis=axes,
+                       keepdims=True)
+        x = safe_divide(x, absmax)
+
+        if self._output_range not in (False, True):  # True = (-1, +1)
+            output_range = self._output_range
+
+            if not self._input_is_positive_only:
+                x = (x+1) / 2
+            x = K.clip(x, 0, 1)
+
+            x = output_range[0] + (x * (output_range[1]-output_range[0]))
+        else:
+            x = K.clip(x, -1, 1)
+
+        return x
+
+
+class Print(_Map):
+    def _apply_map(self, x):
+        return K.print_tensor(x)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class Greater(keras.layers.Layer):
+    def call(self, x):
+        a, b = x
+        return K.greater(a, b)
+
+
+class Less(keras.layers.Layer):
+    def call(self, x):
+        a, b = x
+        return K.less(a, b)
+
+
+class GreaterThanZero(keras.layers.Layer):
+    def call(self, x):
+        return K.greater(x, K.constant(0))
+
+
+class LessThanZero(keras.layers.Layer):
+    def call(self, x):
+        return K.less(x, K.constant(0))
+
+
+class GreaterEqual(keras.layers.Layer):
+    def call(self, x):
+        a, b = x
+        return K.greater_equal(a, b)
+
+
+class LessEqual(keras.layers.Layer):
+    def call(self, x):
+        a, b = x
+        return K.less_equal(a, b)
+
+
+class GreaterEqualThanZero(keras.layers.Layer):
+    def call(self, x):
+        return K.greater_equal(x, K.constant(0))
+
+
+class LessEqualThanZero(keras.layers.Layer):
+    def call(self, x):
+        return K.less_equal(x, K.constant(0))
+
+
+class Transpose(keras.layers.Layer):
+
+    def __init__(self, axes=None, **kwargs):
+        self._axes = axes
+        super(Transpose, self).__init__(**kwargs)
+
+    def call(self, x):
+        if self._axes is None:
+            return K.transpose(x)
+        else:
+            return K.permute_dimensions(x, self._axes)
+
+    def compute_output_shape(self, input_shape):
+        if self._axes is None:
+            return input_shape[::-1]
+        else:
+            return tuple(np.asarray(input_shape)[list(self._axes)])
+
+
+class Dot(keras.layers.Layer):
+
+    def call(self, x):
+        a, b = x
+        return K.dot(a, b)
+
+    def compute_output_shape(self, input_shapes):
+        return (input_shapes[0][0], input_shapes[1][1])
+
+
+class Divide(keras.layers.Layer):
+
+    def call(self, x):
+        a, b = x
+        return a / b
+
+    def compute_output_shape(self, input_shapes):
+        return input_shapes[0]
+
+
+class SafeDivide(keras.layers.Layer):
+
+    def __init__(self, *args, **kwargs):
+        factor = kwargs.pop("factor", None)
+        if factor is None:
+            factor = K.epsilon()
+        self._factor = factor
+
+        return super(SafeDivide, self).__init__(*args, **kwargs)
+
+    def call(self, x):
+        a, b = x
+        return a / (b + iK.to_floatx(K.equal(b, K.constant(0))) * self._factor)
+
+    def compute_output_shape(self, input_shapes):
+        return input_shapes[0]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class Repeat(keras.layers.Layer):
+
+    def __init__(self, n, axis, *args, **kwargs):
+        self._n = n
+        self._axis = axis
+        return super(Repeat, self).__init__(*args, **kwargs)
+
+    def call(self, x):
+        return K.repeat_elements(x, self._n, self._axis)
+
+    def compute_output_shape(self, input_shapes):
+        if isinstance(input_shapes, list):
+            input_shape = input_shapes[0]
+        else:
+            input_shape = input_shapes
+
+        if input_shape[0] is None:
+            return input_shape
+        else:
+            return (input_shape[0]*self._n,)+input_shape[1:]
+
+
+class Reshape(keras.layers.Layer):
+
+    def __init__(self, shape, *args, **kwargs):
+        self._shape = shape
+        return super(Reshape, self).__init__(*args, **kwargs)
+
+    def call(self, x):
+        return K.reshape(x, self._shape)
+
+    def compute_output_shape(self, input_shapes):
+        return tuple(x if x >= 0 else None for x in self._shape)
+
+
+class MultiplyWithLinspace(keras.layers.Layer):
+
+    def __init__(self, start, end, n=1, axis=-1, *args, **kwargs):
+        self._start = start
+        self._end = end
+        self._n = n
+        self._axis = axis
+        return super(MultiplyWithLinspace, self).__init__(*args, **kwargs)
+
+    def call(self, x):
+        linspace = (self._start +
+                    (self._end-self._start) *
+                    (K.arange(self._n, dtype=K.floatx())/self._n))
+
+        # Make broadcastable.
+        shape = np.ones(len(K.int_shape(x)))
+        shape[self._axis] = self._n
+        linspace = K.reshape(linspace, shape)
+        return x * linspace
+
+    def compute_output_shape(self, input_shapes):
+        ret = input_shapes[:]
+        ret = (ret[:self._axis] +
+               (max(self._n, ret[self._axis]),) +
+               ret[self._axis+1:])
+        return ret
+
+
+class TestPhaseGaussianNoise(keras.layers.GaussianNoise):
+
+    def call(self, inputs):
+        # Always add Gaussian noise!
+        return super(TestPhaseGaussianNoise, self).call(inputs, training=True)
+
+
+class ExtractConv2DPatches(keras.layers.Layer):
+
+    def __init__(self,
+                 kernel_shape,
+                 depth,
+                 strides,
+                 rates,
+                 padding,
+                 *args,
+                 **kwargs):
+        self._kernel_shape = kernel_shape
+        self._depth = depth
+        self._strides = strides
+        self._rates = rates
+        self._padding = padding
+        return super(ExtractConv2DPatches, self).__init__(*args, **kwargs)
+
+    def call(self, x):
+        return iK.extract_conv2d_patches(x,
+                                         self._kernel_shape,
+                                         self._strides,
+                                         self._rates,
+                                         self._padding)
+
+    def compute_output_shape(self, input_shapes):
+        if K.image_data_format() == 'channels_first':
+            space = input_shapes[2:]
+            new_space = []
+            for i in range(len(space)):
+                new_dim = conv_utils.conv_output_length(
+                    space[i],
+                    self._kernel_shape[i],
+                    padding=self._padding,
+                    stride=self._strides[i],
+                    dilation=self._rates[i])
+                new_space.append(new_dim)
+
+        if K.image_data_format() == 'channels_last':
+            space = input_shapes[1:-1]
+            new_space = []
+            for i in range(len(space)):
+                new_dim = conv_utils.conv_output_length(
+                    space[i],
+                    self._kernel_shape[i],
+                    padding=self._padding,
+                    stride=self._strides[i],
+                    dilation=self._rates[i])
+                new_space.append(new_dim)
+
+        return ((input_shapes[0],) +
+                tuple(new_space) +
+                (np.product(self._kernel_shape) * self._depth,))
+
+
+class RunningMeans(keras.layers.Layer):
+
+    def __init__(self, *args, **kwargs):
+        self.stateful = True
+        super(RunningMeans, self).__init__(*args, **kwargs)
+
+    def build(self, input_shapes):
+        means_shape, counts_shape = input_shapes
+
+        self.means = self.add_weight(shape=means_shape,
+                                     initializer="zeros",
+                                     name="means",
+                                     trainable=False)
+        self.counts = self.add_weight(shape=counts_shape,
+                                      initializer="zeros",
+                                      name="counts",
+                                      trainable=False)
+        self.built = True
+
+    def call(self, x):
+        def safe_divide(a, b):
+            return a / (b + iK.to_floatx(K.equal(b, K.constant(0))) * 1)
+
+        means, counts = x
+
+        new_counts = counts + self.counts
+
+        # If new_means are not used for the model output,
+        # the following part of the code will be executed after
+        # self.counts is updated, therefore we cannot use it
+        # hereafter.
+        factor_new = safe_divide(counts, new_counts)
+        factor_old = K.ones_like(factor_new) - factor_new
+        new_means = self.means * factor_old + means * factor_new
+
+        # Update state.
+        self.add_update([
+            K.update(self.means, new_means),
+            K.update(self.counts, new_counts),
+        ])
+
+        return [new_means, new_counts]
+
+    def compute_output_shape(self, input_shapes):
+        return input_shapes
+
+
+class Broadcast(keras.layers.Layer):
+
+    def call(self, x):
+        target_shapped, x = x
+        return target_shapped * 0 + x
+
+    def compute_output_shape(self, input_shapes):
+        return input_shapes[0]
+
+
+class Gather(keras.layers.Layer):
+
+    def call(self, inputs):
+        x, index = inputs
+        return iK.gather(x, 1, index)
+
+    def compute_output_shape(self, input_shapes):
+        return (input_shapes[0][0], input_shapes[1][0])+input_shapes[0][2:]
+
+
+class GatherND(keras.layers.Layer):
+
+    def call(self, inputs):
+        x, indices = inputs
+        return iK.gather_nd(x, indices)
+
+    def compute_output_shape(self, input_shapes):
+        return input_shapes[1][:2]+input_shapes[0][2:]

+ 0 - 0
original_model/innvestigate/tests/__init__.py


+ 0 - 0
original_model/innvestigate/tests/analyzer/__init__.py


+ 231 - 0
original_model/innvestigate/tests/analyzer/test_base.py

@@ -0,0 +1,231 @@
+# Get Python six functionality:
+from __future__ import \
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import pytest
+
+
+from innvestigate.utils.tests import dryrun
+
+from innvestigate.analyzer import BaselineGradient
+from innvestigate.analyzer import Gradient
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__BasicGraphReversal():
+
+    def method1(model):
+        return BaselineGradient(model)
+
+    def method2(model):
+        return Gradient(model)
+
+    dryrun.test_equal_analyzer(method1,
+                               method2,
+                               "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__BasicGraphReversal():
+
+    def method1(model):
+        return BaselineGradient(model)
+
+    def method2(model):
+        return Gradient(model)
+
+    dryrun.test_equal_analyzer(method1,
+                               method2,
+                               "mnist.*")
+
+
+# @pytest.mark.fast
+# @pytest.mark.precommit
+# def test_fast__ContainerGraphReversal():
+
+#     def method1(model):
+#         return Gradient(model)
+
+#     def method2(model):
+#         Create container execution
+#         model = keras.models.Model(inputs=model.inputs,
+#                                    outputs=model(model.inputs))
+#         return Gradient(model)
+
+#     dryrun.test_equal_analyzer(method1,
+#                                method2,
+#                                "trivia.*:mnist.log_reg")
+
+
+# @pytest.mark.precommit
+# def test_precommit__ContainerGraphReversal():
+
+#     def method1(model):
+#         return Gradient(model)
+
+#     def method2(model):
+#         Create container execution
+#         model = keras.models.Model(inputs=model.inputs,
+#                                    outputs=model(model.inputs))
+#         return Gradient(model)
+
+#     dryrun.test_equal_analyzer(method1,
+#                                method2,
+#                                "mnist.*")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__AnalyzerNetworkBase_neuron_selection_max():
+
+    def method(model):
+        return Gradient(model, neuron_selection_mode="max_activation")
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__AnalyzerNetworkBase_neuron_selection_max():
+
+    def method(model):
+        return Gradient(model, neuron_selection_mode="max_activation")
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__AnalyzerNetworkBase_neuron_selection_index():
+
+    class CustomAnalyzer(Gradient):
+
+        def analyze(self, X):
+            index = 0
+            return super(CustomAnalyzer, self).analyze(X, index)
+
+    def method(model):
+        return CustomAnalyzer(model, neuron_selection_mode="index")
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__AnalyzerNetworkBase_neuron_selection_index():
+
+    class CustomAnalyzer(Gradient):
+
+        def analyze(self, X):
+            index = 3
+            return super(CustomAnalyzer, self).analyze(X, index)
+
+    def method(model):
+        return CustomAnalyzer(model, neuron_selection_mode="index")
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__BaseReverseNetwork_reverse_debug():
+
+    def method(model):
+        return Gradient(model, reverse_verbose=True)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__BaseReverseNetwork_reverse_debug():
+
+    def method(model):
+        return Gradient(model, reverse_verbose=True)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__BaseReverseNetwork_reverse_check_minmax():
+
+    def method(model):
+        return Gradient(model, reverse_verbose=True,
+                        reverse_check_min_max_values=True)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__BaseReverseNetwork_reverse_check_minmax():
+
+    def method(model):
+        return Gradient(model, reverse_verbose=True,
+                        reverse_check_min_max_values=True)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__BaseReverseNetwork_reverse_check_finite():
+
+    def method(model):
+        return Gradient(model, reverse_verbose=True, reverse_check_finite=True)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__BaseReverseNetwork_reverse_check_finite():
+
+    def method(model):
+        return Gradient(model, reverse_verbose=True, reverse_check_finite=True)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__SerializeAnalyzerBase():
+
+    def method(model):
+        return BaselineGradient(model)
+
+    dryrun.test_serialize_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__SerializeReverseAnalyzerkBase():
+
+    def method(model):
+        return Gradient(model)
+
+    dryrun.test_serialize_analyzer(method, "trivia.*:mnist.log_reg")
+ 

+ 241 - 0
original_model/innvestigate/tests/analyzer/test_deeplift.py

@@ -0,0 +1,241 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.layers
+import keras.models
+import numpy as np
+import pytest
+try:
+    import deeplift
+except ImportError:
+    deeplift = None
+
+from innvestigate.utils.tests import dryrun
+
+from innvestigate.analyzer import DeepLIFT
+from innvestigate.analyzer import DeepLIFTWrapper
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__DeepLIFT():
+
+    def method(model):
+        return DeepLIFT(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__DeepLIFT():
+
+    def method(model):
+        return DeepLIFT(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.precommit
+def test_precommit__DeepLIFT_Rescale():
+
+    def method(model):
+        if keras.backend.image_data_format() == "channels_first":
+            input_shape = (1, 28, 28)
+        else:
+            input_shape = (28, 28, 1)
+        model = keras.models.Sequential([
+            keras.layers.Dense(10, input_shape=input_shape),
+            keras.layers.ReLU(),
+        ])
+        return DeepLIFT(model)
+
+    dryrun.test_analyzer(method, "mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__DeepLIFT_neuron_selection_index():
+
+    class CustomAnalyzer(DeepLIFT):
+
+        def analyze(self, X):
+            index = 0
+            return super(CustomAnalyzer, self).analyze(X, index)
+
+    def method(model):
+        return CustomAnalyzer(model, neuron_selection_mode="index")
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.precommit
+def test_precommit__DeepLIFT_larger_batch_size():
+
+    class CustomAnalyzer(DeepLIFT):
+
+        def analyze(self, X):
+            X = np.concatenate((X, X), axis=0)
+            return super(CustomAnalyzer, self).analyze(X)[0:1]
+
+    def method(model):
+        return CustomAnalyzer(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.skip("There is a design issue to be fixed.")
+@pytest.mark.precommit
+def test_precommit__DeepLIFT_larger_batch_size_with_index():
+
+    class CustomAnalyzer(DeepLIFT):
+
+        def analyze(self, X):
+            index = 0
+            X = np.concatenate((X, X), axis=0)
+            return super(CustomAnalyzer, self).analyze(X, index)[0:1]
+
+    def method(model):
+        return CustomAnalyzer(model, neuron_selection_mode="index")
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__DeepLIFT():
+
+    def method(model):
+        return DeepLIFT(model)
+
+    dryrun.test_analyzer(method, "imagenet.*")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+require_deeplift = pytest.mark.skipif(deeplift is None,
+                                      reason="Package deeplift is required.")
+
+
+@require_deeplift
+@pytest.mark.fast
+@pytest.mark.precommit
+@pytest.mark.skip(reason="DeepLIFT does not work with skip connection.")
+def test_fast__DeepLIFTWrapper():
+
+    def method(model):
+        return DeepLIFTWrapper(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@require_deeplift
+@pytest.mark.precommit
+def test_precommit__DeepLIFTWrapper():
+
+    def method(model):
+        return DeepLIFTWrapper(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@require_deeplift
+@pytest.mark.precommit
+def test_precommit__DeepLIFTWrapper_neuron_selection_index():
+
+    class CustomAnalyzer(DeepLIFTWrapper):
+
+        def analyze(self, X):
+            index = 0
+            return super(CustomAnalyzer, self).analyze(X, index)
+
+    def method(model):
+        return CustomAnalyzer(model, neuron_selection_mode="index")
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@require_deeplift
+@pytest.mark.precommit
+def test_precommit__DeepLIFTWrapper_larger_batch_size():
+
+    class CustomAnalyzer(DeepLIFTWrapper):
+
+        def analyze(self, X):
+            X = np.concatenate((X, X), axis=0)
+            return super(CustomAnalyzer, self).analyze(X)[0:1]
+
+    def method(model):
+        return CustomAnalyzer(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@require_deeplift
+@pytest.mark.precommit
+def test_precommit__DeepLIFTWrapper_larger_batch_size_with_index():
+
+    class CustomAnalyzer(DeepLIFTWrapper):
+
+        def analyze(self, X):
+            index = 0
+            X = np.concatenate((X, X), axis=0)
+            return super(CustomAnalyzer, self).analyze(X, index)[0:1]
+
+    def method(model):
+        return CustomAnalyzer(model, neuron_selection_mode="index")
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@require_deeplift
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__DeepLIFTWrapper():
+
+    def method(model):
+        return DeepLIFTWrapper(model)
+
+    dryrun.test_analyzer(method, "imagenet.*")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__DeepLIFT_serialize():
+
+    def method(model):
+        return DeepLIFT(model)
+
+    dryrun.test_serialize_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__DeepLIFTWrapper_serialize():
+
+    def method(model):
+        return DeepLIFTWrapper(model)
+
+    with pytest.raises(AssertionError):
+        # Issue in deeplift.
+        dryrun.test_serialize_analyzer(method, "trivia.*:mnist.log_reg")

+ 82 - 0
original_model/innvestigate/tests/analyzer/test_deeptaylor.py

@@ -0,0 +1,82 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import pytest
+
+
+from innvestigate.utils.tests import dryrun
+
+from innvestigate.analyzer import DeepTaylor
+from innvestigate.analyzer import BoundedDeepTaylor
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__DeepTaylor():
+
+    def method(model):
+        return DeepTaylor(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__DeepTaylor():
+
+    def method(model):
+        return DeepTaylor(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__DeepTaylor():
+
+    def method(model):
+        return DeepTaylor(model)
+
+    dryrun.test_analyzer(method, "imagenet.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__BoundedDeepTaylor():
+
+    def method(model):
+        return BoundedDeepTaylor(model, low=-1, high=1)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__BoundedDeepTaylor():
+
+    def method(model):
+        return BoundedDeepTaylor(model, low=-1, high=1)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__BoundedDeepTaylor():
+
+    def method(model):
+        return BoundedDeepTaylor(model, low=-1, high=1)
+
+    dryrun.test_analyzer(method, "imagenet.*")

+ 337 - 0
original_model/innvestigate/tests/analyzer/test_gradient_based.py

@@ -0,0 +1,337 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import pytest
+
+
+from innvestigate.utils.tests import dryrun
+
+from innvestigate.analyzer import BaselineGradient
+from innvestigate.analyzer import Gradient
+
+from innvestigate.analyzer import InputTimesGradient
+
+from innvestigate.analyzer import Deconvnet
+from innvestigate.analyzer import GuidedBackprop
+
+from innvestigate.analyzer import IntegratedGradients
+
+from innvestigate.analyzer import SmoothGrad
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__BaselineGradient():
+
+    def method(model):
+        return BaselineGradient(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__BaselineGradient():
+
+    def method(model):
+        return BaselineGradient(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__BaselineGradient():
+
+    def method(model):
+        return BaselineGradient(model)
+
+    dryrun.test_analyzer(method, "imagenet.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__BaselineGradient_pp_None():
+
+    def method(model):
+        return BaselineGradient(model, postprocess=None)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__BaselineGradient_pp_None():
+
+    def method(model):
+        return BaselineGradient(model, postprocess=None)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__BaselineGradient_pp_square():
+
+    def method(model):
+        return BaselineGradient(model, postprocess="square")
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__BaselineGradient_pp_square():
+
+    def method(model):
+        return BaselineGradient(model, postprocess="square")
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__Gradient():
+
+    def method(model):
+        return Gradient(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__Gradient():
+
+    def method(model):
+        return Gradient(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__Gradient():
+
+    def method(model):
+        return Gradient(model)
+
+    dryrun.test_analyzer(method, "imagenet.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__Gradient_pp_None():
+
+    def method(model):
+        return Gradient(model, postprocess=None)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__Gradient_pp_None():
+
+    def method(model):
+        return Gradient(model, postprocess=None)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__Gradient_pp_square():
+
+    def method(model):
+        return Gradient(model, postprocess="square")
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__Gradient_pp_square():
+
+    def method(model):
+        return Gradient(model, postprocess="square")
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__InputTimesGradient():
+
+    def method(model):
+        return InputTimesGradient(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__InputTimesGradient():
+
+    def method(model):
+        return InputTimesGradient(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__InputTimesGradient():
+
+    def method(model):
+        return InputTimesGradient(model)
+
+    dryrun.test_analyzer(method, "imagenet.*")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__Deconvnet():
+
+    def method(model):
+        return Deconvnet(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__Deconvnet():
+
+    def method(model):
+        return Deconvnet(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__Deconvnet():
+
+    def method(model):
+        return Deconvnet(model)
+
+    dryrun.test_analyzer(method, "imagenet.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__GuidedBackprop():
+
+    def method(model):
+        return GuidedBackprop(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__GuidedBackprop():
+
+    def method(model):
+        return GuidedBackprop(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__GuidedBackprop():
+
+    def method(model):
+        return GuidedBackprop(model)
+
+    dryrun.test_analyzer(method, "imagenet.*")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__IntegratedGradients():
+
+    def method(model):
+        return IntegratedGradients(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__IntegratedGradients():
+
+    def method(model):
+        return IntegratedGradients(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__IntegratedGradients():
+
+    def method(model):
+        return IntegratedGradients(model, steps=2)
+
+    dryrun.test_analyzer(method, "imagenet.*")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__SmoothGrad():
+
+    def method(model):
+        return SmoothGrad(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__SmoothGrad():
+
+    def method(model):
+        return SmoothGrad(model)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__SmoothGrad():
+
+    def method(model):
+        return SmoothGrad(model, augment_by_n=2)
+
+    dryrun.test_analyzer(method, "imagenet.*")

+ 51 - 0
original_model/innvestigate/tests/analyzer/test_init.py

@@ -0,0 +1,51 @@
+# Get Python six functionality:
+from __future__ import \
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import pytest
+
+import keras.layers
+import keras.models
+
+from innvestigate import create_analyzer
+from innvestigate.analyzer import analyzers
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__create_analyzers():
+
+    fake_model = keras.models.Sequential([
+        keras.layers.Dense(10, input_shape=(10,))
+    ])
+    for name in analyzers.keys():
+        try:
+            create_analyzer(name, fake_model)
+        except KeyError:
+            # Name should be found!
+            raise
+        except:
+            # Some analyzers require parameters...
+            pass
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__create_analyzers_wrong_name():
+
+    fake_model = keras.models.Sequential([
+        keras.layers.Dense(10, input_shape=(10,))
+    ])
+    with pytest.raises(KeyError):
+        create_analyzer("wrong name", fake_model)

+ 57 - 0
original_model/innvestigate/tests/analyzer/test_misc.py

@@ -0,0 +1,57 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import pytest
+
+
+from innvestigate.utils.tests import dryrun
+
+from innvestigate.analyzer import Input
+from innvestigate.analyzer import Random
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__Input():
+
+    def method(model):
+        return Input(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__Random():
+
+    def method(model):
+        return Random(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__SerializeRandom():
+
+    def method(model):
+        return Random(model)
+
+    dryrun.test_serialize_analyzer(method, "trivia.*:mnist.log_reg")

+ 134 - 0
original_model/innvestigate/tests/analyzer/test_pattern_based.py

@@ -0,0 +1,134 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import pytest
+
+
+from innvestigate.utils.tests import dryrun
+
+from innvestigate.analyzer import PatternNet
+from innvestigate.analyzer import PatternAttribution
+
+
+# todo: add again a traint/test case for mnist
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__PatternNet():
+
+    def method(model):
+        # enough for test purposes, only pattern application is tested here
+        # pattern computation is tested separately.
+        # assume that one dim weights are biases, drop them.
+        patterns = [x for x in model.get_weights()
+                    if len(x.shape) > 1]
+        return PatternNet(model, patterns=patterns)
+
+    dryrun.test_analyzer(method, "mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__PatternNet():
+
+    def method(model):
+        # enough for test purposes, only pattern application is tested here
+        # pattern computation is tested separately.
+        # assume that one dim weights are biases, drop them.
+        patterns = [x for x in model.get_weights()
+                    if len(x.shape) > 1]
+        return PatternNet(model, patterns=patterns)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__PatternNet():
+
+    def method(model):
+        # enough for test purposes, only pattern application is tested here
+        # pattern computation is tested separately.
+        # assume that one dim weights are biases, drop them.
+        patterns = [x for x in model.get_weights()
+                    if len(x.shape) > 1]
+        return PatternNet(model, patterns=patterns)
+
+    dryrun.test_analyzer(method, "imagenet.vgg16:imagenet.vgg19")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__PatternAttribution():
+
+    def method(model):
+        # enough for test purposes, only pattern application is tested here
+        # pattern computation is tested separately.
+        # assume that one dim weights are biases, drop them.
+        patterns = [x for x in model.get_weights()
+                    if len(x.shape) > 1]
+        return PatternAttribution(model, patterns=patterns)
+
+    dryrun.test_analyzer(method, "mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__PatternAttribution():
+
+    def method(model):
+        # enough for test purposes, only pattern application is tested here
+        # pattern computation is tested separately.
+        # assume that one dim weights are biases, drop them.
+        patterns = [x for x in model.get_weights()
+                    if len(x.shape) > 1]
+        return PatternAttribution(model, patterns=patterns)
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.slow
+@pytest.mark.application
+@pytest.mark.imagenet
+def test_imagenet__PatternAttribution():
+
+    def method(model):
+        # enough for test purposes, only pattern application is tested here
+        # pattern computation is tested separately.
+        # assume that one dim weights are biases, drop them.
+        patterns = [x for x in model.get_weights()
+                    if len(x.shape) > 1]
+        return PatternAttribution(model, patterns=patterns)
+
+    dryrun.test_analyzer(method, "imagenet.vgg16:imagenet.vgg19")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__SerializePatternNet():
+
+    def method(model):
+        # enough for test purposes, only pattern application is tested here
+        # pattern computation is tested separately.
+        # assume that one dim weights are biases, drop them.
+        patterns = [x for x in model.get_weights()
+                    if len(x.shape) > 1]
+        return PatternNet(model, patterns=patterns)
+
+    dryrun.test_serialize_analyzer(method, "mnist.log_reg")

+ 242 - 0
original_model/innvestigate/tests/analyzer/test_relevance_based.py

@@ -0,0 +1,242 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import pytest
+
+
+from innvestigate.utils.tests import dryrun
+
+from innvestigate.analyzer import BaselineLRPZ
+from innvestigate.analyzer import LRPZ
+from innvestigate.analyzer import LRPZIgnoreBias
+from innvestigate.analyzer import LRPZPlus
+from innvestigate.analyzer import LRPZPlusFast
+from innvestigate.analyzer import LRPEpsilon
+from innvestigate.analyzer import LRPEpsilonIgnoreBias
+from innvestigate.analyzer import LRPWSquare
+from innvestigate.analyzer import LRPFlat
+from innvestigate.analyzer import LRPAlpha2Beta1
+from innvestigate.analyzer import LRPAlpha2Beta1IgnoreBias
+from innvestigate.analyzer import LRPAlpha1Beta0
+from innvestigate.analyzer import LRPAlpha1Beta0IgnoreBias
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__BaselineLRPZ():
+
+    def method(model):
+        return BaselineLRPZ(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPZ():
+
+    def method(model):
+        return LRPZ(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_fast__LRPZ_resnet50():
+
+    def method(model):
+        return LRPZ(model)
+
+    dryrun.test_analyzer(method, "imagenet.resnet50")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPZ__equal_BaselineLRPZ():
+
+    def method1(model):
+        return BaselineLRPZ(model)
+
+    def method2(model):
+        # LRP-Z with bias
+        return LRPZ(model)
+
+    dryrun.test_equal_analyzer(method1,
+                               method2,
+                               # mind this only works for
+                               # networks with relu, max,
+                               # activations and no
+                               # skip connections!
+                               "trivia.dot:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPZ__with_input_layer_rule():
+
+    def method(model):
+        return LRPZ(model, input_layer_rule="Flat")
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPZ__with_boxed_input_layer_rule():
+
+    def method(model):
+        return LRPZ(model, input_layer_rule=(-10, 10))
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPZIgnoreBias():
+
+    def method(model):
+        return LRPZIgnoreBias(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPZPlus():
+
+    def method(model):
+        return LRPZPlus(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPZPlusFast():
+
+    def method(model):
+        return LRPZPlusFast(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPEpsilon():
+
+    def method(model):
+        return LRPEpsilon(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPEpsilonIgnoreBias():
+
+    def method(model):
+        return LRPEpsilonIgnoreBias(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPWSquare():
+
+    def method(model):
+        return LRPWSquare(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPFlat():
+
+    def method(model):
+        return LRPFlat(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPAlpha2Beta1():
+
+    def method(model):
+        return LRPAlpha2Beta1(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPAlpha2Beta1IgnoreBias():
+
+    def method(model):
+        return LRPAlpha2Beta1IgnoreBias(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPAlpha1Beta0():
+
+    def method(model):
+        return LRPAlpha1Beta0(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__LRPAlpha1Beta0IgnoreBias():
+
+    def method(model):
+        return LRPAlpha1Beta0IgnoreBias(model)
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__SerializeLRPZ():
+
+    def method(model):
+        return LRPZ(model)
+
+    dryrun.test_serialize_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__SerializeLRPAlpha2Beta1():
+
+    def method(model):
+        return LRPAlpha2Beta1(model)
+
+    dryrun.test_serialize_analyzer(method, "trivia.*:mnist.log_reg")

+ 157 - 0
original_model/innvestigate/tests/analyzer/test_wrapper.py

@@ -0,0 +1,157 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import pytest
+
+
+from innvestigate.utils.tests import dryrun
+
+from innvestigate.analyzer import WrapperBase
+from innvestigate.analyzer import AugmentReduceBase
+from innvestigate.analyzer import GaussianSmoother
+from innvestigate.analyzer import PathIntegrator
+
+from innvestigate.analyzer import Gradient
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__WrapperBase():
+
+    def method(model):
+        return WrapperBase(Gradient(model))
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__WrapperBase():
+
+    def method(model):
+        return WrapperBase(Gradient(model))
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__SerializeWrapperBase():
+
+    def method(model):
+        return WrapperBase(Gradient(model))
+
+    dryrun.test_serialize_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__AugmentReduceBase():
+
+    def method(model):
+        return AugmentReduceBase(Gradient(model))
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__AugmentReduceBase():
+
+    def method(model):
+        return AugmentReduceBase(Gradient(model))
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__SerializeAugmentReduceBase():
+
+    def method(model):
+        return AugmentReduceBase(Gradient(model))
+
+    dryrun.test_serialize_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__GaussianSmoother():
+
+    def method(model):
+        return GaussianSmoother(Gradient(model))
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__GaussianSmoother():
+
+    def method(model):
+        return GaussianSmoother(Gradient(model))
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__SerializeGaussianSmoother():
+
+    def method(model):
+        return GaussianSmoother(Gradient(model))
+
+    dryrun.test_serialize_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__PathIntegrator():
+
+    def method(model):
+        return PathIntegrator(Gradient(model))
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__PathIntegrator():
+
+    def method(model):
+        return PathIntegrator(Gradient(model))
+
+    dryrun.test_analyzer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__SerializePathIntegrator():
+
+    def method(model):
+        return PathIntegrator(Gradient(model))
+
+    dryrun.test_serialize_analyzer(method, "trivia.*:mnist.log_reg")

+ 0 - 0
original_model/innvestigate/tests/tools/__init__.py


+ 382 - 0
original_model/innvestigate/tests/tools/test_pattern.py

@@ -0,0 +1,382 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import pytest
+
+
+from keras.datasets import mnist
+import keras.layers
+import keras.models
+from keras.models import Model
+import keras.optimizers
+import numpy as np
+import unittest
+
+from innvestigate.utils.tests import dryrun
+
+import innvestigate
+from innvestigate.tools import PatternComputer
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__PatternComputer_dummy_parallel():
+
+    def method(model):
+        return PatternComputer(model, pattern_type="dummy",
+                               compute_layers_in_parallel=True)
+
+    dryrun.test_pattern_computer(method, "mnist.log_reg")
+
+
+@pytest.mark.skip("Feature not supported.")
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__PatternComputer_dummy_sequential():
+
+    def method(model):
+        return PatternComputer(model, pattern_type="dummy",
+                               compute_layers_in_parallel=False)
+
+    dryrun.test_pattern_computer(method, "mnist.log_reg")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__PatternComputer_linear():
+
+    def method(model):
+        return PatternComputer(model, pattern_type="linear")
+
+    dryrun.test_pattern_computer(method, "mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__PatternComputer_linear():
+
+    def method(model):
+        return PatternComputer(model, pattern_type="linear")
+
+    dryrun.test_pattern_computer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__PatternComputer_relupositive():
+
+    def method(model):
+        return PatternComputer(model, pattern_type="relu.positive")
+
+    dryrun.test_pattern_computer(method, "mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__PatternComputer_relupositive():
+
+    def method(model):
+        return PatternComputer(model, pattern_type="relu.positive")
+
+    dryrun.test_pattern_computer(method, "mnist.*")
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__PatternComputer_relunegative():
+
+    def method(model):
+        return PatternComputer(model, pattern_type="relu.negative")
+
+    dryrun.test_pattern_computer(method, "mnist.log_reg")
+
+
+@pytest.mark.precommit
+def test_precommit__PatternComputer_relunegative():
+
+    def method(model):
+        return PatternComputer(model, pattern_type="relu.negative")
+
+    dryrun.test_pattern_computer(method, "mnist.*")
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+class HaufePatternExample(unittest.TestCase):
+
+    def test(self):
+        np.random.seed(234354346)
+        # need many samples to get close to optimum and stable numbers
+        n = 1000
+
+        a_s = np.asarray([1, 0]).reshape((1, 2))
+        a_d = np.asarray([1, 1]).reshape((1, 2))
+        y = np.random.uniform(size=(n, 1))
+        eps = np.random.rand(n, 1)
+
+        X = y * a_s + eps * a_d
+
+        model = keras.models.Sequential(
+            [keras.layers.Dense(1, input_shape=(2,), use_bias=True), ]
+        )
+        model.compile(optimizer=keras.optimizers.Adam(lr=1), loss="mse")
+        model.fit(X, y, epochs=20, verbose=0).history
+        self.assertTrue(model.evaluate(X, y, verbose=0) < 0.05)
+
+        pc = PatternComputer(model, pattern_type="linear")
+        A = pc.compute(X)[0]
+        W = model.get_weights()[0]
+
+        #print(a_d, model.get_weights()[0])
+        #print(a_s, A)
+
+        def allclose(a, b):
+            return np.allclose(a, b, rtol=0.05, atol=0.05)
+
+        # perpendicular to a_d
+        self.assertTrue(allclose(a_d.ravel(), abs(W.ravel())))
+        # estimated pattern close to true pattern
+        self.assertTrue(allclose(a_s.ravel(), A.ravel()))
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def fetch_data():
+    # the data, shuffled and split between train and test sets
+    (x_train, y_train), (x_test, y_test) = mnist.load_data()
+
+    x_train = (x_train.reshape(60000, 1, 28, 28) - 127.5) / 127.5
+    x_test = (x_test.reshape(10000, 1, 28, 28) - 127.5) / 127.5
+    x_train = x_train.astype('float32')
+    x_test = x_test.astype('float32')
+
+    return x_train[:100], y_train[:100], x_test[:10], y_test[:10]
+
+
+def create_model(clazz):
+    num_classes = 10
+
+    network = clazz(
+        (None, 1, 28, 28),
+        num_classes,
+        dense_units=1024,
+        dropout_rate=0.25)
+    model_wo_sm = Model(inputs=network["in"], outputs=network["out"])
+    model_w_sm = Model(inputs=network["in"], outputs=network["sm_out"])
+    return model_wo_sm, model_w_sm
+
+
+def train_model(model, data, epochs=20):
+    batch_size = 128
+    num_classes = 10
+
+    x_train, y_train, x_test, y_test = data
+    # convert class vectors to binary class matrices
+    y_train = keras.utils.to_categorical(y_train, num_classes)
+    y_test = keras.utils.to_categorical(y_test, num_classes)
+
+    model.compile(loss='categorical_crossentropy',
+                  optimizer=keras.optimizers.RMSprop(),
+                  metrics=['accuracy'])
+
+    model.fit(x_train, y_train,
+              batch_size=batch_size,
+              epochs=epochs,
+              verbose=0)
+    model.evaluate(x_test, y_test, batch_size=batch_size, verbose=0)
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+class MnistPatternExample_dense_linear(unittest.TestCase):
+
+    def test(self):
+        np.random.seed(234354346)
+        model_class = innvestigate.utils.tests.networks.base.mlp_2dense
+
+        data = fetch_data()
+        model, modelp = create_model(model_class)
+        train_model(modelp, data, epochs=10)
+        model.set_weights(modelp.get_weights())
+
+        analyzer = innvestigate.create_analyzer("pattern.net", model,
+                                                pattern_type="linear")
+        analyzer.fit(data[0], batch_size=256, verbose=0)
+
+        patterns = analyzer._patterns
+        W = model.get_weights()[0]
+        W2D = W.reshape((-1, W.shape[-1]))
+        X = data[0].reshape((data[0].shape[0], -1))
+        Y = np.dot(X, W2D)
+
+        def safe_divide(a, b):
+            return a / (b + (b == 0))
+
+        mean_x = X.mean(axis=0)
+        mean_y = Y.mean(axis=0)
+        mean_xy = np.dot(X.T, Y) / Y.shape[0]
+        ExEy = mean_x[:, None] * mean_y[None, :]
+        cov_xy = mean_xy - ExEy
+        w_cov_xy = np.diag(np.dot(W2D.T, cov_xy))
+        A = safe_divide(cov_xy, w_cov_xy[None, :])
+
+        def allclose(a, b):
+            return np.allclose(a, b, rtol=0.05, atol=0.05)
+        #print(A.sum(), patterns[0].sum())
+        self.assertTrue(allclose(A.ravel(), patterns[0].ravel()))
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+class MnistPatternExample_dense_relu(unittest.TestCase):
+
+    def test(self):
+        np.random.seed(234354346)
+        model_class = innvestigate.utils.tests.networks.base.mlp_2dense
+
+        data = fetch_data()
+        model, modelp = create_model(model_class)
+        train_model(modelp, data, epochs=10)
+        model.set_weights(modelp.get_weights())
+
+        analyzer = innvestigate.create_analyzer("pattern.net", model,
+                                                pattern_type="relu")
+        analyzer.fit(data[0], batch_size=256, verbose=0)
+        patterns = analyzer._patterns
+        W, b = model.get_weights()[:2]
+        W2D = W.reshape((-1, W.shape[-1]))
+        X = data[0].reshape((data[0].shape[0], -1))
+        Y = np.dot(X, W2D)
+
+        mask = np.dot(X, W2D) + b > 0
+        count = mask.sum(axis=0)
+
+        def safe_divide(a, b):
+            return a / (b + (b == 0))
+
+        mean_x = safe_divide(np.dot(X.T, mask), count)
+        mean_y = Y.mean(axis=0)
+        mean_xy = safe_divide(np.dot(X.T, Y * mask), count)
+
+        ExEy = mean_x * mean_y
+
+        cov_xy = mean_xy - ExEy
+        w_cov_xy = np.diag(np.dot(W2D.T, cov_xy))
+        A = safe_divide(cov_xy, w_cov_xy[None, :])
+
+        def allclose(a, b):
+            return np.allclose(a, b, rtol=0.05, atol=0.05)
+        #print(A.sum(), patterns[0].sum())
+        self.assertTrue(allclose(A.ravel(), patterns[0].ravel()))
+
+
+# def extract_2d_patches(X, conv_layer):
+
+#     X_in = X
+#     kernel_shape = conv_layer.kernel_size
+#     strides = conv_layer.strides
+#     rates = conv_layer.dilation_rate
+#     padding = conv_layer.padding
+
+#     assert all([x == 1 for x in rates])
+#     assert all([x == 3 for x in kernel_shape])
+#     assert all([x == 1 for x in strides])
+
+#     if padding.lower() == "same":
+#         tmp = np.ones(list(X.shape[:2])+[x+3 for x in X.shape[2:]],
+#                       dtype=X.dtype)
+#         tmp[:, :, 1:-2, 1:-2] = X
+#         X = tmp
+
+#     out_shape = [int(np.ceil((x-k)/s))
+#                  for x, k, s in zip(X.shape[2:], kernel_shape, strides)]
+#     n_patches = np.prod(list(X.shape[:2])+out_shape)
+#     dimensions = X.shape[1]*kernel_shape[0]*kernel_shape[1]
+#     ret = np.empty((n_patches, dimensions), dtype=X.dtype)
+
+#     i_ret = 0
+#     for j in range(X.shape[2]-kernel_shape[0]):
+#         for k in range(X.shape[3]-kernel_shape[1]):
+#             patches = X[:, :, j:j+kernel_shape[0], k:k+kernel_shape[1]]
+#             patches = patches.reshape((-1, dimensions))
+#             ret[i_ret:i_ret+X.shape[0]] = patches
+#             i_ret += X.shape[0]
+
+#     if True:
+#         import tensorflow as tf
+#         with tf.Session():
+#             tf_ret = tf.extract_image_patches(
+#                 images=X_in.transpose((0, 2, 3, 1)),
+#                 ksizes=[1, kernel_shape[0], kernel_shape[1], 1],
+#                 strides=[1, strides[0], strides[1], 1],
+#                 rates=[1, rates[0], rates[1], 1],
+#                 padding=padding.upper()).eval()
+
+#         tf_ret = tf_ret.reshape((-1, tf_ret.shape[-1]))
+#         #print(tf_ret.shape, ret.shape)
+#         assert tf_ret.shape == ret.shape
+#         #print(tf_ret.mean(), ret.mean())
+#         assert tf_ret.mean() == ret.mean()
+#     assert i_ret == n_patches
+#     return ret
+
+
+# class __disabled__MnistPatternExample_conv_linear(unittest.TestCase):
+
+#     def test(self):
+#         np.random.seed(234354346)
+#         K.set_image_data_format("channels_first")
+#         model_class = innvestigate.utils.tests.networks.base.cnn_2convb_2dense
+#         data = fetch_data()
+#         model, modelp = create_model(model_class)
+#         train_model(modelp, data, epochs=1)
+#         model.set_weights(modelp.get_weights())
+#         analyzer = innvestigate.create_analyzer("pattern.net", model)
+#         analyzer.fit(data[0], pattern_type="linear",
+#                      batch_size=256, verbose=0)
+
+#         patterns = analyzer._patterns
+#         W = model.get_weights()[0]
+#         W2D = W.reshape((-1, W.shape[-1]))
+#         X = extract_2d_patches(data[0], model.layers[1])
+#         Y = np.dot(X, W2D)
+
+#         def safe_divide(a, b):
+#             return a / (b + (b == 0))
+
+#         mean_x = X.mean(axis=0)
+#         mean_y = Y.mean(axis=0)
+#         mean_xy = np.dot(X.T, Y) / Y.shape[0]
+
+#         ExEy = mean_x[:, None] * mean_y[None, :]
+#         cov_xy = mean_xy - ExEy
+#         w_cov_xy = np.diag(np.dot(W2D.T, cov_xy))
+#         A = safe_divide(cov_xy, w_cov_xy[None, :])
+
+#         def allclose(a, b):
+#             return np.allclose(a, b, rtol=0.05, atol=0.05)
+
+#         self.assertTrue(allclose(A.ravel(), patterns[0].ravel()))

+ 102 - 0
original_model/innvestigate/tests/tools/test_perturbate.py

@@ -0,0 +1,102 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.layers
+import keras.models
+import numpy as np
+import pytest
+
+import innvestigate.tools.perturbate
+import innvestigate.utils as iutils
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__PerturbationAnalysis():
+    # Some test data
+    if keras.backend.image_data_format() == "channels_first":
+        input_shape = (2, 1, 4, 4)
+    else:
+        input_shape = (2, 4, 4, 1)
+    x = np.arange(2 * 4 * 4).reshape(input_shape)
+    generator = iutils.BatchSequence([x, np.zeros(x.shape[0])], batch_size=x.shape[0])
+
+    # Simple model
+    model = keras.models.Sequential([
+            keras.layers.Flatten(input_shape=x.shape[1:]),
+            keras.layers.Dense(1, use_bias=False),
+    ])
+
+    weights = np.arange(4 * 4 * 1).reshape((4 * 4, 1))
+    model.layers[-1].set_weights([weights])
+    model.compile(loss='mean_squared_error', optimizer='sgd')
+
+    expected_output = np.array([[1240.], [3160.]])
+    assert np.all(np.isclose(model.predict(x), expected_output))
+
+    # Analyzer
+    analyzer = innvestigate.create_analyzer("gradient",
+                                              model,
+                                              postprocess="abs")
+
+    # Run perturbation analysis
+    perturbation = innvestigate.tools.perturbate.Perturbation("zeros", region_shape=(2, 2), in_place=False)
+
+    perturbation_analysis = innvestigate.tools.perturbate.PerturbationAnalysis(analyzer, model, generator, perturbation, recompute_analysis=False,
+                                                 steps=3, regions_per_step=1, verbose=False)
+
+    scores = perturbation_analysis.compute_perturbation_analysis()
+
+    expected_scores = np.array([5761600.0, 1654564.0, 182672.0, 21284.0])
+    assert np.all(np.isclose(scores, expected_scores))
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__Perturbation():
+    if keras.backend.image_data_format() == "channels_first":
+        input_shape = (1, 1, 4, 4)
+    else:
+        input_shape = (1, 4, 4, 1)
+    x = np.arange(1 * 4 * 4).reshape(input_shape)
+
+    perturbation = innvestigate.tools.perturbate.Perturbation("zeros", region_shape=(2, 2), in_place=False)
+
+    analysis = np.zeros((4, 4))
+    analysis[:2, 2:] = 1
+    analysis[2:, :2] = 2
+    analysis[2:, 2:] = 3
+    analysis = analysis.reshape(input_shape)
+
+    if keras.backend.image_data_format() == "channels_last":
+        x = np.moveaxis(x, 3, 1)
+        analysis = np.moveaxis(analysis, 3, 1)
+
+    analysis = perturbation.reduce_function(analysis, axis=1, keepdims=True)
+
+    aggregated_regions = perturbation.aggregate_regions(analysis)
+    assert np.all(np.isclose(aggregated_regions[0, 0, :, :], np.array([[0, 1], [2, 3]])))
+
+    ranks = perturbation.compute_region_ordering(aggregated_regions)
+    assert np.all(np.isclose(ranks[0, 0, :, :], np.array([[3, 2], [1, 0]])))
+
+    perturbation_mask_regions = perturbation.compute_perturbation_mask(ranks, 1)
+    assert np.all(perturbation_mask_regions == np.array([[0, 0], [0, 1]]))
+
+    perturbation_mask_regions = perturbation.compute_perturbation_mask(ranks, 4)
+    assert np.all(perturbation_mask_regions == np.array([[1, 1], [1, 1]]))
+
+    perturbation_mask_regions = perturbation.compute_perturbation_mask(ranks, 0)
+    assert np.all(perturbation_mask_regions == np.array([[0, 0], [0, 0]]))

+ 0 - 0
original_model/innvestigate/tests/utils/__init__.py


+ 0 - 0
original_model/innvestigate/tests/utils/keras/__init__.py


+ 80 - 0
original_model/innvestigate/tests/utils/keras/test_graph.py

@@ -0,0 +1,80 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.models
+import pytest
+
+
+from innvestigate.utils.keras import graph as kgraph
+from innvestigate.utils.tests import networks
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__get_model_execution_graph():
+
+    network_filter = "trivia.*:mnist.log_reg"
+
+    for network in networks.iterator(network_filter):
+
+        model = keras.models.Model(inputs=network["in"],
+                                   outputs=network["out"])
+
+        graph = kgraph.get_model_execution_graph(model)
+        kgraph.print_model_execution_graph(graph)
+
+
+@pytest.mark.precommit
+def test_commit__get_model_execution_graph():
+
+    network_filter = "mnist.*"
+
+    for network in networks.iterator(network_filter):
+
+        model = keras.models.Model(inputs=network["in"],
+                                   outputs=network["out"])
+
+        graph = kgraph.get_model_execution_graph(model)
+        kgraph.print_model_execution_graph(graph)
+
+
+@pytest.mark.precommit
+def test_precommit__get_model_execution_graph_resnet50():
+
+    network_filter = "imagenet.resnet50"
+
+    for network in networks.iterator(network_filter):
+
+        model = keras.models.Model(inputs=network["in"],
+                                   outputs=network["out"])
+
+        graph = kgraph.get_model_execution_graph(model)
+        kgraph.print_model_execution_graph(graph)
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__get_model_execution_graph_with_inputs():
+
+    network_filter = "trivia.*:mnist.log_reg"
+
+    for network in networks.iterator(network_filter):
+
+        model = keras.models.Model(inputs=network["in"],
+                                   outputs=network["out"])
+
+        graph = kgraph.get_model_execution_graph(model,
+                                                 keep_input_layers=True)
+        kgraph.print_model_execution_graph(graph)

+ 32 - 0
original_model/innvestigate/tests/utils/test_visualizations.py

@@ -0,0 +1,32 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import numpy as np
+import pytest
+
+import innvestigate.utils.visualizations as ivis
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__visualizations():
+    def get_X():
+        return np.random.rand(1, 28, 28, 3)
+
+    ivis.project(get_X())
+    ivis.heatmap(get_X())
+    ivis.graymap(get_X())
+    ivis.gamma(get_X())
+    ivis.clip_quantile(get_X(), 0.95)

+ 0 - 0
original_model/innvestigate/tests/utils/tests/__init__.py


+ 37 - 0
original_model/innvestigate/tests/utils/tests/test_dryrun.py

@@ -0,0 +1,37 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import pytest
+
+
+from innvestigate.utils.tests import dryrun
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__DryRunAnalyzerTestCase():
+    """
+    Sanity test for the TestCase.
+    """
+
+    def method(output_layer):
+
+        class TestAnalyzer(object):
+            def analyze(self, X):
+                return X
+
+        return TestAnalyzer()
+
+    dryrun.test_analyzer(method, "trivia.*:mnist.log_reg")

+ 64 - 0
original_model/innvestigate/tests/utils/tests/test_layer.py

@@ -0,0 +1,64 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.layers
+import numpy as np
+import pytest
+
+
+from innvestigate.analyzer.gradient_based import Gradient
+# Prevent pytest from collecting this class:
+from innvestigate.utils.tests.layer import TestAnalysisHelper as AnalysisHelper
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__TestAnalysisHelper_one_layer():
+
+    layer = keras.layers.Dense(2, input_shape=(3,), use_bias=False)
+    analyzer = Gradient
+    weights = [np.asarray(((1, 2), (3, 4), (5, 6)))]
+
+    helper = AnalysisHelper(layer, analyzer, weights)
+
+    inputs = np.asarray((1, 2, 3))
+    outputs, analysis = helper.run(inputs)
+
+    # Analyzer takes node with max output.
+    i = np.argmax(outputs)
+    gradient = np.dot(weights[0][:, i], np.ones_like(outputs[i]))
+    assert np.allclose(analysis, gradient)
+
+
+@pytest.mark.fast
+@pytest.mark.precommit
+def test_fast__TestAnalysisHelper_two_layers():
+
+    layers = [keras.layers.Dense(2, input_shape=(3,), use_bias=False),
+              keras.layers.Dense(2, use_bias=False)]
+    analyzer = Gradient
+    weights = [np.asarray(((1, 2), (3, 4), (5, 6))),
+               np.asarray(((7, 8), (9, 1)))]
+
+    helper = AnalysisHelper(layers, analyzer, weights)
+
+    inputs = np.asarray((1, 2, 3))
+    outputs, analysis = helper.run(inputs)
+
+    # Analyzer takes node with max output.
+    i = np.argmax(outputs)
+    gradient_middle = np.dot(weights[1][:, i], np.ones_like(outputs[i]))
+    gradient = np.dot(weights[0], gradient_middle)
+    assert np.allclose(analysis, gradient)

+ 9 - 0
original_model/innvestigate/tools/__init__.py

@@ -0,0 +1,9 @@
+
+from .pattern import PatternComputer
+from .perturbate import Perturbation
+from .perturbate import PerturbationAnalysis
+
+# Make pylint ignore the imports
+assert PatternComputer
+assert Perturbation
+assert PerturbationAnalysis

+ 520 - 0
original_model/innvestigate/tools/pattern.py

@@ -0,0 +1,520 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import range
+import six
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.backend as K
+import keras.layers
+import keras.models
+import keras.optimizers
+import keras.utils
+import numpy as np
+
+
+from .. import layers as ilayers
+from .. import utils as iutils
+from ..utils.keras import checks as kchecks
+from ..utils.keras import graph as kgraph
+
+
+__all__ = [
+    "get_active_neuron_io",
+    "get_pattern_class",
+
+    "BasePattern",
+    "DummyPattern",
+    "LinearPattern",
+    "ReLUPositivePattern",
+    "ReLUNegativePattern",
+
+    "PatternComputer",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def get_active_neuron_io(layer, active_node_indices,
+                         return_i=True, return_o=True,
+                         do_activation_search=False):
+    """
+    Returns the neuron-wise input output for the passed layer.
+    This is done while taking care of only considering layer nodes that
+    are listed as active.
+
+    Starting from the passed layer this functions
+    returns the first layer with an activation upstream in the model,
+    if do_activation_search is an execution list.
+    Otherwise the current layer's output is returned.
+    """
+
+    def contains_activation(layer):
+        return (kchecks.contains_activation(layer) and
+                not kchecks.contains_activation(layer, "linear"))
+
+    def get_Xs(node_index):
+        return iutils.to_list(layer.get_input_at(node_index))
+
+    def get_Ys(node_index):
+        ret = iutils.to_list(layer.get_output_at(node_index))
+        if(do_activation_search is not False and
+           not contains_activation(layer)):
+            # Walk along execution graph until we find an activation function,
+            # if current layer has none.
+            execution_list = do_activation_search
+
+            # First find current node.
+            layer_i = None
+            for i, node in enumerate(execution_list):
+                if layer is node[0]:
+                    layer_i = i
+                    break
+
+            assert layer_i is not None
+            assert len(ret) == 1
+            input_to_next_layer = ret[0]
+
+            found = False
+            for i in range(layer_i+1, len(execution_list)):
+                l, Xs, Ys = execution_list[i]
+                if input_to_next_layer in Xs:
+                    if not isinstance(
+                            l,
+                            kchecks.get_activation_search_safe_layers()):
+                        break
+                    if contains_activation(l):
+                        found = Ys
+                        break
+                    assert len(Ys) == 1
+                    input_to_next_layer = Ys[0]
+
+            if found is not False:
+                ret = Ys
+
+        return ret
+
+    # Get neuron-wise io for active layer nodes.
+    tmp = [kgraph.get_layer_neuronwise_io(layer, Xs=get_Xs(i), Ys=get_Ys(i),
+                                          return_i=return_i, return_o=return_o)
+           for i in active_node_indices]
+
+    if len(tmp) == 1:
+        return tmp[0]
+    else:
+        raise NotImplementedError("This code seems not to handle several Ys.")
+        # Layer is applied several times in model.
+        # Concatenate the io of the applications.
+        concatenate = keras.layers.Concatenate(axis=0)
+
+        if return_i and return_o:
+            return (concatenate([x[0] for x in tmp]),
+                    concatenate([x[1] for x in tmp]))
+        else:
+            return concatenate([x[0] for x in tmp])
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class BasePattern(object):
+    """
+    Interface for pattern objects used to compute patterns by the
+    PatternComputer class.
+
+    The basic work-flow is that a pattern computes statistics for the
+    passed layer, which are then used to compute the final pattern.
+    """
+
+    def __init__(self,
+                 model,
+                 layer,
+                 model_tensors=None,
+                 execution_list=None):
+        self.model = model
+        self.layer = layer
+        # All the tensors used by the model.
+        # Allows to filter nodes in layers that do not
+        # belong to this model.
+        self.model_tensors = model_tensors
+        self.execution_list = execution_list
+        self._active_node_indices = self._get_active_node_indices()
+
+    def _get_active_node_indices(self):
+        """
+        A layer can be applied in several models.
+        This functions returns a list with all nodes of the given
+        layer that are active/used in the current model.
+
+        If no model_tensors are passed to the pattern,
+        it is assumed all nodes are active.
+        """
+        n_nodes = kgraph.get_layer_inbound_count(self.layer)
+        if self.model_tensors is None:
+            return list(range(n_nodes))
+        else:
+            ret = []
+            for i in range(n_nodes):
+                output_tensors = iutils.to_list(self.layer.get_output_at(i))
+                # Check if output is used in the model.
+                if all([tmp in self.model_tensors
+                        for tmp in output_tensors]):
+                    ret.append(i)
+            return ret
+
+    def has_pattern(self):
+        return kchecks.contains_kernel(self.layer)
+
+    def stats_from_batch(self):
+        """
+        Creates statistics while the PatternComputer passes the
+        dataset once.
+        """
+        raise NotImplementedError()
+
+    def compute_pattern(self):
+        """
+        Computes the pattern after computing the statistics.
+        """
+        raise NotImplementedError()
+
+
+class DummyPattern(BasePattern):
+    """
+    Computes a dummy pattern for test purposes.
+    """
+
+    def get_stats_from_batch(self):
+        Xs, Ys = get_active_neuron_io(self.layer,
+                                      self._active_node_indices)
+        self.mean_x = ilayers.RunningMeans()
+
+        count = ilayers.CountNonZero(axis=0)(Ys[0])
+        sum_x = ilayers.Dot()([ilayers.Transpose()(Xs[0]), Ys[0]])
+
+        mean_x, count_x = self.mean_x([sum_x, count])
+
+        # Return dummy output to have connected graph!
+        return ilayers.Sum(axis=None)(count_x)
+
+    def compute_pattern(self):
+        return self.mean_x.get_weights()[0]
+
+
+class LinearPattern(BasePattern):
+
+    def _get_neuron_mask(self):
+        """
+        Select which neurons are considered for the pattern computation.
+        """
+        Ys = get_active_neuron_io(self.layer,
+                                  self._active_node_indices,
+                                  return_i=False, return_o=True)
+
+        return ilayers.OnesLike()(Ys[0])
+
+    def get_stats_from_batch(self):
+        # Get the neuron-wise I/O for this layer.
+        layer = kgraph.copy_layer_wo_activation(self.layer,
+                                                keep_bias=False,
+                                                reuse_symbolic_tensors=False)
+        # Readjust the layer nodes.
+        for i in range(kgraph.get_layer_inbound_count(self.layer)):
+            layer(self.layer.get_input_at(i))
+        Xs, Ys = get_active_neuron_io(layer, self._active_node_indices)
+        if len(Ys) != 1:
+            raise ValueError("Assume that kernel layer have only one output.")
+        X, Y = Xs[0], Ys[0]
+
+        # Create layers that keep a running mean for the desired stats.
+        self.mean_x = ilayers.RunningMeans()
+        self.mean_y = ilayers.RunningMeans()
+        self.mean_xy = ilayers.RunningMeans()
+
+        # Compute mask and active neuron counts.
+        mask = ilayers.AsFloatX()(self._get_neuron_mask())
+        Y_masked = keras.layers.multiply([Y, mask])
+        count = ilayers.CountNonZero(axis=0)(mask)
+        count_all = ilayers.Sum(axis=0)(ilayers.OnesLike()(mask))
+
+        # Get means ...
+        def norm(x, count):
+            return ilayers.SafeDivide(factor=1)([x, count])
+
+        # ... along active neurons.
+        mean_x = norm(ilayers.Dot()([ilayers.Transpose()(X), mask]), count)
+        mean_xy = norm(ilayers.Dot()([ilayers.Transpose()(X), Y_masked]),
+                       count)
+
+        _, a = self.mean_x([mean_x, count])
+        _, b = self.mean_xy([mean_xy, count])
+
+        # ... along all neurons.
+        mean_y = norm(ilayers.Sum(axis=0)(Y), count_all)
+        _, c = self.mean_y([mean_y, count_all])
+
+        # Create a dummy output to have a connected graph.
+        # Needs to have the shape (mb_size, 1)
+        dummy = keras.layers.Average()([a, b, c])
+        return ilayers.Sum(axis=None)(dummy)
+
+    def compute_pattern(self):
+        """Computes the patterns according to the formula in the paper."""
+        def safe_divide(a, b):
+            return a / (b + (b == 0))
+
+        W = kgraph.get_kernel(self.layer)
+        W2D = W.reshape((-1, W.shape[-1]))
+
+        mean_x, cnt_x = self.mean_x.get_weights()
+        mean_y, cnt_y = self.mean_y.get_weights()
+        mean_xy, cnt_xy = self.mean_xy.get_weights()
+
+        ExEy = mean_x * mean_y
+        cov_xy = mean_xy - ExEy
+
+        w_cov_xy = np.diag(np.dot(W2D.T, cov_xy))
+        A = safe_divide(cov_xy, w_cov_xy[None, :])
+
+        # update length
+        if False:
+            norm = np.diag(np.dot(W2D.T, A))
+            A = safe_divide(A, norm)
+
+        # check pattern
+        if False:
+            tmp = np.diag(np.dot(W2D.T, A))
+            print("pattern_check", W.shape, tmp.min(), tmp.max())
+
+        return A.reshape(W.shape)
+
+
+class ReLUPositivePattern(LinearPattern):
+
+    def _get_neuron_mask(self):
+        Ys = get_active_neuron_io(self.layer,
+                                  self._active_node_indices,
+                                  return_i=False, return_o=True,
+                                  do_activation_search=self.execution_list)
+        return ilayers.GreaterThanZero()(Ys[0])
+
+
+class ReLUNegativePattern(LinearPattern):
+
+    def _get_neuron_mask(self):
+        Ys = get_active_neuron_io(self.layer,
+                                  self._active_node_indices,
+                                  return_i=False, return_o=True,
+                                  do_activation_search=self.execution_list)
+        return ilayers.LessEqualThanZero()(Ys[0])
+
+
+def get_pattern_class(pattern_type):
+    return {
+        "dummy": DummyPattern,
+
+        "linear": LinearPattern,
+        "relu": ReLUPositivePattern,
+        "relu.positive": ReLUPositivePattern,
+        "relu.negative": ReLUNegativePattern,
+    }.get(pattern_type, pattern_type)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class PatternComputer(object):
+    """Pattern computer.
+
+    Computes a pattern for each layer with a kernel of a given model.
+
+    :param model: A Keras model.
+    :param pattern_type: A string or a tuple of strings. Valid types are
+      'linear', 'relu', 'relu.positive', 'relu.negative'.
+    :param compute_layers_in_parallel: Not supported yet.
+      Compute all patterns at once.
+      Otherwise computer layer after layer.
+    :param gpus: Not supported yet. Gpus to use.
+    """
+
+    def __init__(self, model,
+                 pattern_type="linear",
+                 # todo: this options seems to be buggy,
+                 # if it sequential tensorflow still pushes all models to gpus
+                 compute_layers_in_parallel=True,
+                 gpus=None):
+        self.model = model
+
+        # Break cyclic import.
+        import innvestigate.analyzer.pattern_based
+        supported_layers = (
+            innvestigate.analyzer.pattern_based.SUPPORTED_LAYER_PATTERNNET)
+        for layer in self.model.layers:
+            if not isinstance(layer, supported_layers):
+                raise Exception("Model contains not supported layer: %s"
+                                % layer)
+
+        pattern_types = iutils.to_list(pattern_type)
+        self.pattern_types = {k: get_pattern_class(k)
+                              for k in pattern_types}
+        self.compute_layers_in_parallel = compute_layers_in_parallel
+        self.gpus = gpus
+
+        if self.compute_layers_in_parallel is False:
+            raise NotImplementedError("Not supported.")
+
+    def _create_computers(self):
+        """
+        Creates pattern objects and Keras models that are used to collect
+        statistics and compute patterns.
+
+        We compute the patterns by first computing statistics within
+        the Keras framework, which are then used to compute the patterns.
+
+        This is based on a workaround. We connect the stats computation
+        via dummy outputs to a model's output and then iterate over the
+        dataset to compute statistics.
+        """
+        # Create a broadcasting function that is used to connect
+        # the dummy outputs.
+        # Broadcaster has shape (mini_batch_size, 1)
+        reduce_axes = list(range(len(K.int_shape(self.model.inputs[0]))))[1:]
+        dummy_broadcaster = ilayers.Sum(axis=reduce_axes,
+                                        keepdims=True)(self.model.inputs[0])
+
+        def broadcast(x):
+            return ilayers.Broadcast()([dummy_broadcaster, x])
+
+        # Collect all tensors that are part of a model's execution.
+        layers, execution_list, _ = kgraph.trace_model_execution(self.model)
+        model_tensors = set()
+        for _, input_tensors, output_tensors in execution_list:
+            for t in input_tensors+output_tensors:
+                model_tensors.add(t)
+
+        # Create pattern instances and collect the dummy outputs.
+        self._pattern_instances = {k: [] for k in self.pattern_types}
+        computer_outputs = []
+        for layer_id, layer in enumerate(layers):
+            # This does not work with containers!
+            # They should be replaced by trace_model_execution.
+            if kchecks.is_network(layer):
+                raise Exception("Network in network is not suppored!")
+            for pattern_type, clazz in six.iteritems(self.pattern_types):
+                pinstance = clazz(self.model, layer,
+                                  model_tensors=model_tensors,
+                                  execution_list=execution_list)
+                if pinstance.has_pattern() is False:
+                    continue
+                self._pattern_instances[pattern_type].append(pinstance)
+                dummy_output = pinstance.get_stats_from_batch()
+                # Broadcast dummy_output to right shape.
+                computer_outputs += iutils.to_list(broadcast(dummy_output))
+
+        # Now we create one or more Keras models to train the patterns.
+        self._n_computer_outputs = len(computer_outputs)
+        if self.compute_layers_in_parallel is True:
+            self._computers = [
+                keras.models.Model(inputs=self.model.inputs,
+                                   outputs=computer_outputs)
+            ]
+        else:
+            self._computers = [
+                keras.models.Model(inputs=self.model.inputs,
+                                   outputs=computer_output)
+                for computer_output in computer_outputs
+            ]
+
+        # Distribute computation on more gpus.
+        if self.gpus is not None and self.gpus > 1:
+            raise NotImplementedError("Not supported yet.")
+            self._computers = [keras.utils.multi_gpu_model(tmp, gpus=self.gpus)
+                               for tmp in self._computers]
+
+    def compute(self, X, batch_size=32, verbose=0):
+        """
+        Compute and return the patterns for the model and the data `X`.
+
+        :param X: Data to compute patterns.
+        :param batch_size: Batch size to use.
+        :param verbose: As for keras model.fit.
+        """
+        generator = iutils.BatchSequence(X, batch_size)
+        return self.compute_generator(generator, verbose=verbose)
+
+    def compute_generator(self, generator, **kwargs):
+        """
+        Compute and return the patterns for the model and the data `X`.
+
+        :param generator: Data to compute patterns.
+        :param kwargs: Same as for keras model.fit_generator.
+        """
+        self._create_computers()
+
+        # We don't do gradient updates.
+        class NoOptimizer(keras.optimizers.Optimizer):
+            def get_updates(self, *args, **kwargs):
+                return []
+        optimizer = NoOptimizer()
+        # We only pass the training data once.
+        if "epochs" in kwargs and kwargs["epochs"] != 1:
+            raise ValueError("Pattern are computed with "
+                             "a closed form solution. "
+                             "Only need to do one epoch.")
+        kwargs["epochs"] = 1
+
+        if self.compute_layers_in_parallel is True:
+            n_dummy_outputs = self._n_computer_outputs
+        else:
+            n_dummy_outputs = 1
+
+        # Augment the input with dummy targets.
+        def get_dummy_targets(Xs):
+            n, dtype = Xs[0].shape[0], Xs[0].dtype
+            dummy = np.ones(shape=(n, 1), dtype=dtype)
+            return [dummy for _ in range(n_dummy_outputs)]
+
+        if isinstance(generator, keras.utils.Sequence):
+            generator = iutils.TargetAugmentedSequence(generator,
+                                                       get_dummy_targets)
+        else:
+            base_generator = generator
+
+            def generator(*args, **kwargs):
+                for Xs in base_generator(*args, **kwargs):
+                    Xs = iutils.to_list(Xs)
+                    yield Xs, get_dummy_targets(Xs)
+
+        # Compile models.
+        for computer in self._computers:
+            computer.compile(optimizer=optimizer, loss=lambda x, y: x)
+
+        # Compute pattern statistics.
+        for computer in self._computers:
+            computer.fit_generator(generator, **kwargs)
+
+        # Compute and retrieve the actual patterns.
+        pis = self._pattern_instances
+        patterns = {ptype: [tmp.compute_pattern() for tmp in pis[ptype]]
+                    for ptype in self.pattern_types}
+
+        # Free memory.
+        del self._computers
+        del self._pattern_instances
+
+        if len(self.pattern_types) == 1:
+            return patterns[list(self.pattern_types.keys())[0]]
+        else:
+            return patterns

+ 390 - 0
original_model/innvestigate/tools/perturbate.py

@@ -0,0 +1,390 @@
+# Get Python six functionality:
+from __future__ import \
+    absolute_import, print_function, division, unicode_literals
+from builtins import range
+import six
+
+import numpy as np
+import warnings
+import time
+
+import keras.backend as K
+from keras.utils import Sequence
+from keras.utils.data_utils import OrderedEnqueuer, GeneratorEnqueuer
+
+import innvestigate.utils
+
+
+class Perturbation:
+    """Perturbation of pixels based on analysis result.
+
+    :param perturbation_function: Defines the function with which the samples are perturbated. Can be a function or a string that defines a predefined perturbation function.
+    :type perturbation_function: function or callable or str
+    :param num_perturbed_regions: Number of regions to be perturbed.
+    :type num_perturbed_regions: int
+    :param reduce_function: Function to reduce the analysis result to one channel, e.g. mean or max function.
+    :type reduce_function: function or callable
+    :param aggregation_function: Function to aggregate the analysis over subregions.
+    :type aggregation_function: function or callable
+    :param pad_mode: How to pad if the image cannot be subdivided into an integer number of regions. As in numpy.pad.
+    :type pad_mode: str or function or callable
+    :param in_place: If true, the perturbations are performed in place, i.e. the input samples are modified.
+    :type in_place: bool
+    :param value_range: Minimal and maximal value after perturbation as a tuple: (min_val, max_val). The input is clipped to this range
+    :type value_range: tuple"""
+
+    def __init__(self, perturbation_function, num_perturbed_regions=0, region_shape=(9, 9), reduce_function=np.mean,
+                 aggregation_function=np.mean, pad_mode="reflect", in_place=False, value_range=None):
+        if isinstance(perturbation_function, six.string_types):
+            if perturbation_function == "zeros":
+                # This is equivalent to setting the perturbated values to the channel mean if the data are standardized.
+                self.perturbation_function = np.zeros_like
+            elif perturbation_function == "gaussian":
+                # If scale = 1/3, most of the values will be between -1 and 1
+                self.perturbation_function = lambda x: np.random.normal(loc=0.0, scale=0.3, size=x.shape)
+            elif perturbation_function == "mean":
+                self.perturbation_function = np.mean
+            elif perturbation_function == "invert":
+                self.perturbation_function = lambda x: -x
+            else:
+                raise ValueError("Perturbation function type '{}' not known.".format(perturbation_function))
+        elif callable(perturbation_function):
+            self.perturbation_function = perturbation_function
+        else:
+            raise TypeError("Cannot handle perturbation function of type {}.".format(type(perturbation_function)))
+
+        self.num_perturbed_regions = num_perturbed_regions
+        self.region_shape = region_shape
+        self.reduce_function = reduce_function
+        self.aggregation_function = aggregation_function
+
+        self.pad_mode = pad_mode  # numpy.pad
+
+        self.in_place = in_place
+        self.value_range = value_range
+
+    @staticmethod
+    def compute_perturbation_mask(ranks, num_perturbated_regions):
+        perturbation_mask_regions = ranks <= num_perturbated_regions - 1
+        return perturbation_mask_regions
+
+    @staticmethod
+    def compute_region_ordering(aggregated_regions):
+        # 0 means highest scoring region
+        new_shape = tuple(aggregated_regions.shape[:2]) + (-1,)
+        order = np.argsort(-aggregated_regions.reshape(new_shape), axis=-1)
+        ranks = order.argsort().reshape(aggregated_regions.shape)
+        return ranks
+
+    def expand_regions_to_pixels(self, regions):
+        # Resize to pixels (repeat values).
+        # (n, c, h_aggregated_region, w_aggregated_region) -> (n, c, h_aggregated_region, h_region, w_aggregated_region, w_region)
+        regions_reshaped = np.expand_dims(np.expand_dims(regions, axis=3), axis=5)
+        region_pixels = np.repeat(regions_reshaped, self.region_shape[0], axis=3)
+        region_pixels = np.repeat(region_pixels, self.region_shape[1], axis=5)
+        assert region_pixels.shape[0] == regions.shape[0] and region_pixels.shape[2:] == (
+            regions.shape[2], self.region_shape[0], regions.shape[3], self.region_shape[1]), region_pixels.shape
+
+        return region_pixels
+
+    def reshape_region_pixels(self, region_pixels, target_shape):
+        # Reshape to output shape
+        pixels = region_pixels.reshape(target_shape)
+        assert region_pixels.shape[0] == pixels.shape[0] and region_pixels.shape[1] == pixels.shape[1] and \
+               region_pixels.shape[2] * region_pixels.shape[3] == pixels.shape[2] and region_pixels.shape[4] * \
+               region_pixels.shape[5] == pixels.shape[3]
+        return pixels
+
+    def pad(self, analysis):
+        pad_shape = self.region_shape - np.array(analysis.shape[2:]) % self.region_shape
+        assert np.all(pad_shape < self.region_shape)
+
+        # Pad half the window before and half after (on h and w axes)
+        pad_shape_before = (pad_shape / 2).astype(int)
+        pad_shape_after = pad_shape - pad_shape_before
+        pad_shape = (
+            (0, 0), (0, 0), (pad_shape_before[0], pad_shape_after[0]), (pad_shape_before[1], pad_shape_after[1]))
+        analysis = np.pad(analysis, pad_shape, self.pad_mode)
+        assert np.all(np.array(analysis.shape[2:]) % self.region_shape == 0), analysis.shape[2:]
+        return analysis, pad_shape_before
+
+    def reshape_to_regions(self, analysis):
+        aggregated_shape = tuple((np.array(analysis.shape[2:]) / self.region_shape).astype(int))
+        regions = analysis.reshape(
+            (analysis.shape[0], analysis.shape[1], aggregated_shape[0], self.region_shape[0], aggregated_shape[1],
+             self.region_shape[1]))
+        return regions
+
+    def aggregate_regions(self, analysis):
+        regions = self.reshape_to_regions(analysis)
+        aggregated_regions = self.aggregation_function(regions, axis=(3, 5))
+        return aggregated_regions
+
+    def perturbate_regions(self, x, perturbation_mask_regions):
+        # Perturbate every region in tensor.
+        # A single region (at region_x, region_y in sample) should be in mask[sample, channel, region_x, :, region_y, :]
+
+        x_perturbated = self.reshape_to_regions(x)
+        for sample_idx, channel_idx, region_row, region_col in np.ndindex(perturbation_mask_regions.shape):
+            region = x_perturbated[sample_idx, channel_idx, region_row, :, region_col, :]
+            region_mask = perturbation_mask_regions[sample_idx, channel_idx, region_row, region_col]
+            if region_mask:
+                x_perturbated[sample_idx, channel_idx, region_row, :, region_col, :] = self.perturbation_function(
+                    region)
+
+                if self.value_range is not None:
+                    np.clip(x_perturbated,
+                            self.value_range[0],
+                            self.value_range[1],
+                            x_perturbated)
+        x_perturbated = self.reshape_region_pixels(x_perturbated, x.shape)
+        return x_perturbated
+
+    def perturbate_on_batch(self, x, analysis):
+        """
+        :param x: Batch of images.
+        :type x: numpy.ndarray
+        :param analysis: Analysis of this batch.
+        :type analysis: numpy.ndarray
+        :return: Batch of perturbated images
+        :rtype: numpy.ndarray
+        """
+        if K.image_data_format() == "channels_last":
+            x = np.moveaxis(x, 3, 1)
+            analysis = np.moveaxis(analysis, 3, 1)
+        if not self.in_place:
+            x = np.copy(x)
+        assert analysis.shape == x.shape, analysis.shape
+        original_shape = x.shape
+        # reduce the analysis along channel axis -> n x 1 x h x w
+        analysis = self.reduce_function(analysis, axis=1, keepdims=True)
+        assert analysis.shape == (x.shape[0], 1, x.shape[2], x.shape[3]), analysis.shape
+
+        padding = not np.all(np.array(analysis.shape[2:]) % self.region_shape == 0)
+        if padding:
+            analysis, pad_shape_before_analysis = self.pad(analysis)
+            x, pad_shape_before_x = self.pad(x)
+        aggregated_regions = self.aggregate_regions(analysis)
+
+        # Compute perturbation mask (mask with ones where the input should be perturbated, zeros otherwise)
+        ranks = self.compute_region_ordering(aggregated_regions)
+        perturbation_mask_regions = self.compute_perturbation_mask(ranks, self.num_perturbed_regions)
+        # Perturbate each region
+        x_perturbated = self.perturbate_regions(x, perturbation_mask_regions)
+
+        # Crop the original image region to remove the padding
+        if padding:
+            x_perturbated = x_perturbated[:, :, pad_shape_before_x[0]:pad_shape_before_x[0] + original_shape[2],
+                            pad_shape_before_x[1]:pad_shape_before_x[1] + original_shape[3]]
+
+        if K.image_data_format() == "channels_last":
+            x_perturbated = np.moveaxis(x_perturbated, 1, 3)
+            x = np.moveaxis(x, 1, 3)
+            analysis = np.moveaxis(analysis, 1, 3)
+        return x_perturbated
+
+
+class PerturbationAnalysis:
+    """
+    Performs the perturbation analysis.
+
+    :param analyzer: Analyzer.
+    :type analyzer: innvestigate.analyzer.base.AnalyzerBase
+    :param model: Trained Keras model.
+    :type model: keras.engine.training.Model
+    :param generator: Data generator.
+    :type generator: innvestigate.utils.BatchSequence
+    :param perturbation: Instance of Perturbation class that performs the perturbation.
+    :type perturbation: innvestigate.tools.Perturbation
+    :param steps: Number of perturbation steps.
+    :type steps: int
+    :param regions_per_step: Number of regions that are perturbed per step.
+    :type regions_per_step: float
+    :param recompute_analysis: If true, the analysis is recomputed after each perturbation step.
+    :type recompute_analysis: bool
+    :param verbose: If true, print some useful information, e.g. timing, progress etc.
+    """
+
+    def __init__(self, analyzer, model, generator, perturbation, steps=1, regions_per_step=1, recompute_analysis=False,
+                 verbose=False):
+        self.analyzer = analyzer
+        self.model = model
+        self.generator = generator
+        self.perturbation = perturbation
+        # if not isinstance(perturbation, Perturbation):
+        #     raise TypeError(type(perturbation))
+        self.steps = steps
+        self.regions_per_step = regions_per_step
+        self.recompute_analysis = recompute_analysis
+
+        if not self.recompute_analysis:
+            # Compute the analysis once in the beginning
+            analysis = list()
+            x = list()
+            y = list()
+            for xx, yy in self.generator:
+                x.extend(list(xx))
+                y.extend(list(yy))
+                analysis.extend(list(self.analyzer.analyze(xx)))
+            x = np.array(x)
+            y = np.array(y)
+            analysis = np.array(analysis)
+            self.analysis_generator = innvestigate.utils.BatchSequence([x, y, analysis], batch_size=256)
+        self.verbose = verbose
+
+    def compute_on_batch(self, x, analysis=None, return_analysis=False):
+        """
+        Computes the analysis and perturbes the input batch accordingly.
+
+        :param x: Samples.
+        :param analysis: Analysis of x. If None, it is recomputed.
+        :type x: numpy.ndarray
+        """
+        if analysis is None:
+            analysis = self.analyzer.analyze(x)
+
+        x_perturbated = self.perturbation.perturbate_on_batch(x, analysis)
+        if return_analysis:
+            return x_perturbated, analysis
+        else:
+            return x_perturbated
+
+    def evaluate_on_batch(self, x, y, analysis=None, sample_weight=None):
+        """
+        Perturbs the input batch and scores the model on the perturbed batch.
+
+        :param x: Samples.
+        :type x: numpy.ndarray
+        :param y: Labels.
+        :type y: numpy.ndarray
+        :param analysis: Analysis of x.
+        :type analysis: numpy.ndarray
+        :param sample_weight: Sample weights.
+        :type sample_weight: None
+        :return: List of test scores.
+        :rtype: list
+        """
+        if sample_weight is not None:
+            raise NotImplementedError("Sample weighting is not supported yet.")  # TODO
+        x_perturbated = self.compute_on_batch(x, analysis)
+        score = self.model.test_on_batch(x_perturbated, y, sample_weight=sample_weight)
+        return score
+
+    def evaluate_generator(self, generator, steps=None,
+                           max_queue_size=10,
+                           workers=1,
+                           use_multiprocessing=False):
+        """Evaluates the model on a data generator.
+
+        The generator should return the same kind of data
+        as accepted by `test_on_batch`.
+        For documentation, refer to keras.engine.training.evaluate_generator (https://keras.io/models/model/)
+        """
+
+        steps_done = 0
+        wait_time = 0.01
+        all_outs = []
+        batch_sizes = []
+        is_sequence = isinstance(generator, Sequence)
+        if not is_sequence and use_multiprocessing and workers > 1:
+            warnings.warn(
+                UserWarning('Using a generator with `use_multiprocessing=True`'
+                            ' and multiple workers may duplicate your data.'
+                            ' Please consider using the`keras.utils.Sequence'
+                            ' class.'))
+        if steps is None:
+            if is_sequence:
+                steps = len(generator)
+            else:
+                raise ValueError('`steps=None` is only valid for a generator'
+                                 ' based on the `keras.utils.Sequence` class.'
+                                 ' Please specify `steps` or use the'
+                                 ' `keras.utils.Sequence` class.')
+        enqueuer = None
+
+        try:
+            if workers > 0:
+                if is_sequence:
+                    enqueuer = OrderedEnqueuer(generator,
+                                               use_multiprocessing=use_multiprocessing)
+                else:
+                    enqueuer = GeneratorEnqueuer(generator,
+                                                 use_multiprocessing=use_multiprocessing,
+                                                 wait_time=wait_time)
+                enqueuer.start(workers=workers, max_queue_size=max_queue_size)
+                output_generator = enqueuer.get()
+            else:
+                output_generator = generator
+
+            while steps_done < steps:
+                generator_output = next(output_generator)
+                if not hasattr(generator_output, '__len__'):
+                    raise ValueError('Output of generator should be a tuple '
+                                     '(x, y, sample_weight) '
+                                     'or (x, y). Found: ' +
+                                     str(generator_output))
+                if len(generator_output) == 2:
+                    x, y = generator_output
+                    analysis = None
+                elif len(generator_output) == 3:
+                    x, y, analysis = generator_output
+                else:
+                    raise ValueError('Output of generator should be a tuple '
+                                     '(x, y, analysis) '
+                                     'or (x, y). Found: ' +
+                                     str(generator_output))
+                outs = self.evaluate_on_batch(x, y, analysis=analysis, sample_weight=None)
+
+                if isinstance(x, list):
+                    batch_size = x[0].shape[0]
+                elif isinstance(x, dict):
+                    batch_size = list(x.values())[0].shape[0]
+                else:
+                    batch_size = x.shape[0]
+                if batch_size == 0:
+                    raise ValueError('Received an empty batch. '
+                                     'Batches should at least contain one item.')
+                all_outs.append(outs)
+
+                steps_done += 1
+                batch_sizes.append(batch_size)
+
+        finally:
+            if enqueuer is not None:
+                enqueuer.stop()
+
+        if not isinstance(outs, list):
+            return np.average(np.asarray(all_outs),
+                              weights=batch_sizes)
+        else:
+            averages = []
+            for i in range(len(outs)):
+                averages.append(np.average([out[i] for out in all_outs],
+                                           weights=batch_sizes))
+            return averages
+
+    def compute_perturbation_analysis(self):
+        scores = list()
+        # Evaluate first on original data
+        scores.append(self.model.evaluate_generator(self.generator))
+        self.perturbation.num_perturbed_regions = 1
+        time_start = time.time()
+        for step in range(self.steps):
+            tic = time.time()
+            if self.verbose:
+                print("Step {} of {}: {} regions perturbed.".format(step + 1, self.steps,
+                                                                    self.perturbation.num_perturbed_regions), end=" ")
+            scores.append(self.evaluate_generator(self.analysis_generator))
+            self.perturbation.num_perturbed_regions += self.regions_per_step
+            toc = time.time()
+            if self.verbose:
+                print("Time elapsed: {:.3f} seconds.".format(toc - tic))
+        time_end = time.time()
+
+        if self.verbose:
+            print("Time elapsed for {} steps: {:.3f} seconds.".format(step + 1,
+                                                                      time_end - time_start))  # Use step + 1 instead of self.steps because the analysis can stop prematurely.
+
+        self.perturbation.num_perturbed_regions = 1  # Reset to original value
+        assert len(scores) == self.steps + 1
+        return scores

+ 185 - 0
original_model/innvestigate/utils/__init__.py

@@ -0,0 +1,185 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.backend as K
+import keras.utils
+import math
+
+
+__all__ = [
+    "model_wo_softmax",
+    "to_list",
+
+    "BatchSequence",
+    "TargetAugmentedSequence",
+
+    "preprocess_images",
+    "postprocess_images",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def model_wo_softmax(*args, **kwargs):
+    # Break cyclic import
+    from .keras.graph import model_wo_softmax
+
+    return model_wo_softmax(*args, **kwargs)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def to_list(l):
+    """ If not list, wraps parameter into a list."""
+    if not isinstance(l, list):
+        return [l, ]
+    else:
+        return l
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class BatchSequence(keras.utils.Sequence):
+    """Batch sequence generator.
+
+    Take a (list of) input tensors and a batch size
+    and creates a generators that creates a sequence of batches.
+
+    :param Xs: One or a list of tensors. First axis needs to have same length.
+    :param batch_size: Batch size. Default 32.
+    """
+
+    def __init__(self, Xs, batch_size=32):
+        self.Xs = to_list(Xs)
+        self.single_tensor = len(Xs) == 1
+        self.batch_size = batch_size
+
+        if not self.single_tensor:
+            for X in self.Xs[1:]:
+                assert X.shape[0] == self.Xs[0].shape[0]
+        super(BatchSequence, self).__init__()
+
+    def __len__(self):
+        return int(math.ceil(float(len(self.Xs[0])) / self.batch_size))
+
+    def __getitem__(self, idx):
+        ret = [X[idx*self.batch_size:(idx+1)*self.batch_size]
+               for X in self.Xs]
+
+        if self.single_tensor:
+            return ret[0]
+        else:
+            return tuple(ret)
+
+
+class TargetAugmentedSequence(keras.utils.Sequence):
+    """Augments a sequence with a target on the fly.
+
+    Takes a sequence/generator and a function that
+    creates on the fly for each batch a target.
+    The generator takes a batch from that sequence,
+    computes the target and returns both.
+
+    :param sequence: A sequence or generator.
+    :param augment_f: Takes a batch and returns a target.
+    """
+
+    def __init__(self, sequence, augment_f):
+        self.sequence = sequence
+        self.augment_f = augment_f
+
+        super(TargetAugmentedSequence, self).__init__()
+
+    def __len__(self):
+        return len(self.sequence)
+
+    def __getitem__(self, idx):
+        inputs = self.sequence[idx]
+        if isinstance(inputs, tuple):
+            assert len(inputs) == 1
+            inputs = inputs[0]
+
+        targets = self.augment_f(to_list(inputs))
+        return inputs, targets
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def preprocess_images(images, color_coding=None):
+    """Image preprocessing
+
+    Takes a batch of images and:
+    * Adjust the color axis to the Keras format.
+    * Fixes the color coding.
+
+    :param images: Batch of images with 4 axes.
+    :param color_coding: Determines the color coding.
+      Can be None, 'RGBtoBGR' or 'BGRtoRGB'.
+    :return: The preprocessed batch.
+    """
+
+    ret = images
+    image_data_format = K.image_data_format()
+    # todo: not very general:
+    channels_first = images.shape[1] in [1, 3]
+    if image_data_format == "channels_first" and not channels_first:
+        ret = ret.transpose(0, 3, 1, 2)
+    if image_data_format == "channels_last" and channels_first:
+        ret = ret.transpose(0, 2, 3, 1)
+
+    assert color_coding in [None, "RGBtoBGR", "BGRtoRGB"]
+    if color_coding in ["RGBtoBGR", "BGRtoRGB"]:
+        if image_data_format == "channels_first":
+            ret = ret[:, ::-1, :, :]
+        if image_data_format == "channels_last":
+            ret = ret[:, :, :, ::-1]
+
+    return ret
+
+
+def postprocess_images(images, color_coding=None, channels_first=None):
+    """Image postprocessing
+
+    Takes a batch of images and reverts the preprocessing.
+
+    :param images: A batch of images with 4 axes.
+    :param color_coding: The initial color coding,
+      see :func:`preprocess_images`.
+    :param channels_first: The output channel format.
+    :return: The postprocessed images.
+    """
+
+    ret = images
+    image_data_format = K.image_data_format()
+    assert color_coding in [None, "RGBtoBGR", "BGRtoRGB"]
+    if color_coding in ["RGBtoBGR", "BGRtoRGB"]:
+        if image_data_format == "channels_first":
+            ret = ret[:, ::-1, :, :]
+        if image_data_format == "channels_last":
+            ret = ret[:, :, :, ::-1]
+
+    if image_data_format == "channels_first" and not channels_first:
+        ret = ret.transpose(0, 2, 3, 1)
+    if image_data_format == "channels_last" and channels_first:
+        ret = ret.transpose(0, 3, 1, 2)
+
+    return ret

+ 77 - 0
original_model/innvestigate/utils/keras/__init__.py

@@ -0,0 +1,77 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import zip
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.backend as K
+import numpy as np
+
+
+from ... import utils as iutils
+
+
+__all__ = [
+    "apply",
+    "broadcast_np_tensors_to_keras_tensors",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def apply(layer, inputs):
+    """
+    Apply a layer to input[s].
+
+    A flexible apply that tries to fit input to layers expected input.
+    This is useful when one doesn't know if a layer expects a single tensor
+    or many.
+
+    :param layer: A Keras layer instance.
+    :param inputs: A list of input tensors or a single tensor.
+    """
+
+    if isinstance(inputs, list) and len(inputs) > 1:
+        try:
+            ret = layer(inputs)
+        except (TypeError, AttributeError):
+            # layer expects a single tensor.
+            if len(inputs) != 1:
+                raise ValueError("Layer expects only a single input!")
+            ret = layer(inputs[0])
+    else:
+        ret = layer(inputs[0])
+
+    return iutils.to_list(ret)
+
+
+def broadcast_np_tensors_to_keras_tensors(keras_tensors, np_tensors):
+    """Broadcasts numpy tensors to the shape of Keras tensors.
+
+    :param keras_tensors: The Keras tensors with the target shapes.
+    :param np_tensors: Numpy tensors that should be broadcasted.
+    :return: The broadcasted Numpy tensors.
+    """
+
+    def none_to_one(tmp):
+        return [1 if x is None else x for x in tmp]
+
+    keras_tensors = iutils.to_list(keras_tensors)
+
+    if isinstance(np_tensors, list):
+        ret = [np.broadcast_to(ri, none_to_one(K.int_shape(x)))
+               for x, ri in zip(keras_tensors, np_tensors)]
+    else:
+        ret = [np.broadcast_to(np_tensors,
+                               none_to_one(K.int_shape(x)))
+               for x in keras_tensors]
+
+    return ret

+ 171 - 0
original_model/innvestigate/utils/keras/backend.py

@@ -0,0 +1,171 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import zip
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.backend as K
+
+
+__all__ = [
+    "to_floatx",
+    "gradients",
+    "is_not_finite",
+    "extract_conv2d_patches",
+    "gather",
+    "gather_nd",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def to_floatx(x):
+    return K.cast(x, K.floatx())
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def gradients(Xs, Ys, known_Ys):
+    """Partial derivatives
+
+    Computes the partial derivatives between Ys and Xs and
+    using the gradients for Ys known_Ys.
+
+    :param Xs: List of input tensors.
+    :param Ys: List of output tensors that depend on Xs.
+    :param known_Ys: Gradients for Ys.
+    :return: Gradients for Xs given known_Ys
+    """
+    backend = K.backend()
+    if backend == "theano":
+        # no global import => do not break if module is not present
+        assert len(Ys) == 1
+        import theano.gradient
+        known_Ys = {k: v for k, v in zip(Ys, known_Ys)}
+        # todo: check the stop gradient issue here!
+        return theano.gradient.grad(K.sum(Ys[0]), Xs, known_grads=known_Ys)
+    elif backend == "tensorflow":
+        # no global import => do not break if module is not present
+        import tensorflow
+        return tensorflow.gradients(Ys, Xs,
+                                    grad_ys=known_Ys,
+                                    stop_gradients=Xs)
+    else:
+        # todo: add cntk
+        raise NotImplementedError()
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def is_not_finite(x):
+    """Checks if tensor x is finite, if not throws an exception."""
+
+    backend = K.backend()
+    if backend == "theano":
+        # no global import => do not break if module is not present
+        import theano.tensor
+        return theano.tensor.or_(theano.tensor.isnan(x),
+                                 theano.tensor.isinf(x))
+    elif backend == "tensorflow":
+        # no global import => do not break if module is not present
+        import tensorflow
+        #x = tensorflow.check_numerics(x, "innvestigate - is_finite check")
+        return tensorflow.logical_not(tensorflow.is_finite(x))
+    else:
+        # todo: add cntk
+        raise NotImplementedError()
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def extract_conv2d_patches(x, kernel_shape, strides, rates, padding):
+    """Extracts conv2d patches like TF function extract_image_patches.
+
+    :param x: Input image.
+    :param kernel_shape: Shape of the Keras conv2d kernel.
+    :param strides: Strides of the Keras conv2d layer.
+    :param rates: Dilation rates of the Keras conv2d layer.
+    :param padding: Paddings of the Keras conv2d layer.
+    :return: The extracted patches.
+    """
+
+    backend = K.backend()
+    if backend == "theano":
+        # todo: add theano function.
+        raise NotImplementedError()
+    elif backend == "tensorflow":
+        # no global import => do not break if module is not present
+        import tensorflow
+
+        if K.image_data_format() == "channels_first":
+            x = K.permute_dimensions(x, (0, 2, 3, 1))
+        kernel_shape = [1, kernel_shape[0], kernel_shape[1], 1]
+        strides = [1, strides[0], strides[1], 1]
+        rates = [1, rates[0], rates[1], 1]
+        ret = tensorflow.extract_image_patches(x,
+                                               kernel_shape,
+                                               strides,
+                                               rates,
+                                               padding.upper())
+
+        if K.image_data_format() == "channels_first":
+            # todo: check if we need to permute again.xs
+            pass
+        return ret
+    else:
+        # todo: add cntk
+        raise NotImplementedError()
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def gather(x, axis, indices):
+    """Works as TensorFlow's gather."""
+    backend = K.backend()
+    if backend == "theano":
+        # todo: add theano function.
+        raise NotImplementedError()
+    elif backend == "tensorflow":
+        # no global import => do not break if module is not present
+        import tensorflow
+
+        return tensorflow.gather(x, indices, axis=axis)
+    else:
+        # todo: add cntk
+        raise NotImplementedError()
+
+
+def gather_nd(x, indices):
+    """Works as TensorFlow's gather_nd."""
+    backend = K.backend()
+    if backend == "theano":
+        # todo: add theano function.
+        raise NotImplementedError()
+    elif backend == "tensorflow":
+        # no global import => do not break if module is not present
+        import tensorflow
+
+        return tensorflow.gather_nd(x, indices)
+    else:
+        # todo: add cntk
+        raise NotImplementedError()

+ 435 - 0
original_model/innvestigate/utils/keras/checks.py

@@ -0,0 +1,435 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import inspect
+import keras.engine.topology
+import keras.layers
+import keras.layers.advanced_activations
+import keras.layers.convolutional
+import keras.layers.convolutional_recurrent
+import keras.layers.core
+import keras.layers.cudnn_recurrent
+import keras.layers.embeddings
+import keras.layers.local
+import keras.layers.noise
+import keras.layers.normalization
+import keras.layers.pooling
+import keras.layers.recurrent
+import keras.layers.wrappers
+import keras.legacy.layers
+
+
+# Prevents circular imports.
+def get_kgraph():
+    from . import graph as kgraph
+    return kgraph
+
+
+__all__ = [
+    "get_current_layers",
+    "get_known_layers",
+    "get_activation_search_safe_layers",
+
+    "contains_activation",
+    "contains_kernel",
+    "only_relu_activation",
+    "is_network",
+    "is_convnet_layer",
+    "is_relu_convnet_layer",
+    "is_average_pooling",
+    "is_max_pooling",
+    "is_input_layer",
+    "is_batch_normalization_layer",
+    "is_embedding_layer"
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def get_current_layers():
+    """
+    Returns a list of currently available layers in Keras.
+    """
+    class_set = set([(getattr(keras.layers, name), name)
+                     for name in dir(keras.layers)
+                     if (inspect.isclass(getattr(keras.layers, name)) and
+                         issubclass(getattr(keras.layers, name),
+                                    keras.engine.topology.Layer))])
+    return [x[1] for x in sorted((str(x[0]), x[1]) for x in class_set)]
+
+
+def get_known_layers():
+    """
+    Returns a list of keras layer we are aware of.
+    """
+
+    # Inside function to not break import if Keras changes.
+    KNOWN_LAYERS = (
+        keras.engine.topology.InputLayer,
+        keras.layers.advanced_activations.ELU,
+        keras.layers.advanced_activations.LeakyReLU,
+        keras.layers.advanced_activations.PReLU,
+        keras.layers.advanced_activations.Softmax,
+        keras.layers.advanced_activations.ThresholdedReLU,
+        keras.layers.convolutional.Conv1D,
+        keras.layers.convolutional.Conv2D,
+        keras.layers.convolutional.Conv2DTranspose,
+        keras.layers.convolutional.Conv3D,
+        keras.layers.convolutional.Conv3DTranspose,
+        keras.layers.convolutional.Cropping1D,
+        keras.layers.convolutional.Cropping2D,
+        keras.layers.convolutional.Cropping3D,
+        keras.layers.convolutional.SeparableConv1D,
+        keras.layers.convolutional.SeparableConv2D,
+        keras.layers.convolutional.UpSampling1D,
+        keras.layers.convolutional.UpSampling2D,
+        keras.layers.convolutional.UpSampling3D,
+        keras.layers.convolutional.ZeroPadding1D,
+        keras.layers.convolutional.ZeroPadding2D,
+        keras.layers.convolutional.ZeroPadding3D,
+        keras.layers.convolutional_recurrent.ConvLSTM2D,
+        keras.layers.convolutional_recurrent.ConvRecurrent2D,
+        keras.layers.core.Activation,
+        keras.layers.core.ActivityRegularization,
+        keras.layers.core.Dense,
+        keras.layers.core.Dropout,
+        keras.layers.core.Flatten,
+        keras.layers.core.Lambda,
+        keras.layers.core.Masking,
+        keras.layers.core.Permute,
+        keras.layers.core.RepeatVector,
+        keras.layers.core.Reshape,
+        keras.layers.core.SpatialDropout1D,
+        keras.layers.core.SpatialDropout2D,
+        keras.layers.core.SpatialDropout3D,
+        keras.layers.cudnn_recurrent.CuDNNGRU,
+        keras.layers.cudnn_recurrent.CuDNNLSTM,
+        keras.layers.embeddings.Embedding,
+        keras.layers.local.LocallyConnected1D,
+        keras.layers.local.LocallyConnected2D,
+        keras.layers.Add,
+        keras.layers.Average,
+        keras.layers.Concatenate,
+        keras.layers.Dot,
+        keras.layers.Maximum,
+        keras.layers.Minimum,
+        keras.layers.Multiply,
+        keras.layers.Subtract,
+        keras.layers.noise.AlphaDropout,
+        keras.layers.noise.GaussianDropout,
+        keras.layers.noise.GaussianNoise,
+        keras.layers.normalization.BatchNormalization,
+        keras.layers.pooling.AveragePooling1D,
+        keras.layers.pooling.AveragePooling2D,
+        keras.layers.pooling.AveragePooling3D,
+        keras.layers.pooling.GlobalAveragePooling1D,
+        keras.layers.pooling.GlobalAveragePooling2D,
+        keras.layers.pooling.GlobalAveragePooling3D,
+        keras.layers.pooling.GlobalMaxPooling1D,
+        keras.layers.pooling.GlobalMaxPooling2D,
+        keras.layers.pooling.GlobalMaxPooling3D,
+        keras.layers.pooling.MaxPooling1D,
+        keras.layers.pooling.MaxPooling2D,
+        keras.layers.pooling.MaxPooling3D,
+        keras.layers.recurrent.GRU,
+        keras.layers.recurrent.GRUCell,
+        keras.layers.recurrent.LSTM,
+        keras.layers.recurrent.LSTMCell,
+        keras.layers.recurrent.RNN,
+        keras.layers.recurrent.SimpleRNN,
+        keras.layers.recurrent.SimpleRNNCell,
+        keras.layers.recurrent.StackedRNNCells,
+        keras.layers.wrappers.Bidirectional,
+        keras.layers.wrappers.TimeDistributed,
+        keras.layers.wrappers.Wrapper,
+        keras.legacy.layers.Highway,
+        keras.legacy.layers.MaxoutDense,
+        keras.legacy.layers.Merge,
+        keras.legacy.layers.Recurrent,
+    )
+    return KNOWN_LAYERS
+
+
+def get_activation_search_safe_layers():
+    """
+    Returns a list of keras layer that we can walk along
+    in an activation search.
+    """
+
+    # Inside function to not break import if Keras changes.
+    ACTIVATION_SEARCH_SAFE_LAYERS = (
+        keras.layers.advanced_activations.ELU,
+        keras.layers.advanced_activations.LeakyReLU,
+        keras.layers.advanced_activations.PReLU,
+        keras.layers.advanced_activations.Softmax,
+        keras.layers.advanced_activations.ThresholdedReLU,
+        keras.layers.core.Activation,
+        keras.layers.core.ActivityRegularization,
+        keras.layers.core.Dropout,
+        keras.layers.core.Flatten,
+        keras.layers.core.Reshape,
+        keras.layers.Add,
+        keras.layers.noise.GaussianNoise,
+        keras.layers.normalization.BatchNormalization,
+    )
+    return ACTIVATION_SEARCH_SAFE_LAYERS
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def contains_activation(layer, activation=None):
+    """
+    Check whether the layer contains an activation function.
+    activation is None then we only check if layer can contain an activation.
+    """
+
+    # todo: add test and check this more throughroughly.
+    # rely on Keras convention.
+    if hasattr(layer, "activation"):
+        if activation is not None:
+            return layer.activation == keras.activations.get(activation)
+        else:
+            return True
+    elif isinstance(layer, keras.layers.ReLU):
+        if activation is not None:
+            return (keras.activations.get("relu") ==
+                    keras.activations.get(activation))
+        else:
+            return True
+    elif isinstance(layer, (
+            keras.layers.advanced_activations.ELU,
+            keras.layers.advanced_activations.LeakyReLU,
+            keras.layers.advanced_activations.PReLU,
+            keras.layers.advanced_activations.Softmax,
+            keras.layers.advanced_activations.ThresholdedReLU)):
+        if activation is not None:
+            raise Exception("Cannot detect activation type.")
+        else:
+            return True
+    else:
+        return False
+
+
+def contains_kernel(layer):
+    """
+    Check whether the layer contains a kernel.
+    """
+
+    # TODO: add test and check this more throughroughly.
+    # rely on Keras convention.
+    if hasattr(layer, "kernel") or hasattr(layer, "depthwise_kernel") or hasattr(layer, "pointwise_kernel"):
+        return True
+    else:
+        return False
+
+
+def contains_bias(layer):
+    """
+    Check whether the layer contains a bias.
+    """
+
+    # todo: add test and check this more throughroughly.
+    # rely on Keras convention.
+    if hasattr(layer, "bias"):
+        return True
+    else:
+        return False
+
+
+def only_relu_activation(layer):
+    """Checks if layer contains no or only a ReLU activation."""
+    return (not contains_activation(layer) or
+            contains_activation(layer, None) or
+            contains_activation(layer, "linear") or
+            contains_activation(layer, "relu"))
+
+
+def is_network(layer):
+    """
+    Is network in network?
+    """
+    return isinstance(layer, keras.engine.topology.Network)
+
+
+def is_conv_layer(layer, *args, **kwargs):
+    """Checks if layer is a convolutional layer."""
+    CONV_LAYERS = (
+        keras.layers.convolutional.Conv1D,
+        keras.layers.convolutional.Conv2D,
+        keras.layers.convolutional.Conv2DTranspose,
+        keras.layers.convolutional.Conv3D,
+        keras.layers.convolutional.Conv3DTranspose,
+        keras.layers.convolutional.SeparableConv1D,
+        keras.layers.convolutional.SeparableConv2D,
+        keras.layers.convolutional.DepthwiseConv2D
+    )
+    return isinstance(layer, CONV_LAYERS)
+
+def is_embedding_layer(layer, *args, **kwargs):
+    return isinstance(layer, keras.layers.Embedding)
+
+def is_batch_normalization_layer(layer, *args, **kwargs):
+    """Checks if layer is a batchnorm layer."""
+    return isinstance(layer, keras.layers.normalization.BatchNormalization)
+
+
+def is_add_layer(layer, *args, **kwargs):
+    """Checks if layer is an addition-merge layer."""
+    return isinstance(layer, keras.layers.Add)
+
+
+def is_dense_layer(layer, *args, **kwargs):
+    """Checks if layer is a dense layer."""
+    return isinstance(layer, keras.layers.core.Dense)
+
+
+def is_convnet_layer(layer):
+    """Checks if layer is from a convolutional network."""
+    # Inside function to not break import if Keras changes.
+    CONVNET_LAYERS = (
+        keras.engine.topology.InputLayer,
+        keras.layers.advanced_activations.ELU,
+        keras.layers.advanced_activations.LeakyReLU,
+        keras.layers.advanced_activations.PReLU,
+        keras.layers.advanced_activations.Softmax,
+        keras.layers.advanced_activations.ThresholdedReLU,
+        keras.layers.convolutional.Conv1D,
+        keras.layers.convolutional.Conv2D,
+        keras.layers.convolutional.Conv2DTranspose,
+        keras.layers.convolutional.Conv3D,
+        keras.layers.convolutional.Conv3DTranspose,
+        keras.layers.convolutional.Cropping1D,
+        keras.layers.convolutional.Cropping2D,
+        keras.layers.convolutional.Cropping3D,
+        keras.layers.convolutional.SeparableConv1D,
+        keras.layers.convolutional.SeparableConv2D,
+        keras.layers.convolutional.UpSampling1D,
+        keras.layers.convolutional.UpSampling2D,
+        keras.layers.convolutional.UpSampling3D,
+        keras.layers.convolutional.ZeroPadding1D,
+        keras.layers.convolutional.ZeroPadding2D,
+        keras.layers.convolutional.ZeroPadding3D,
+        keras.layers.core.Activation,
+        keras.layers.core.ActivityRegularization,
+        keras.layers.core.Dense,
+        keras.layers.core.Dropout,
+        keras.layers.core.Flatten,
+        keras.layers.core.Lambda,
+        keras.layers.core.Masking,
+        keras.layers.core.Permute,
+        keras.layers.core.RepeatVector,
+        keras.layers.core.Reshape,
+        keras.layers.core.SpatialDropout1D,
+        keras.layers.core.SpatialDropout2D,
+        keras.layers.core.SpatialDropout3D,
+        keras.layers.embeddings.Embedding,
+        keras.layers.local.LocallyConnected1D,
+        keras.layers.local.LocallyConnected2D,
+        keras.layers.Add,
+        keras.layers.Average,
+        keras.layers.Concatenate,
+        keras.layers.Dot,
+        keras.layers.Maximum,
+        keras.layers.Minimum,
+        keras.layers.Multiply,
+        keras.layers.Subtract,
+        keras.layers.noise.AlphaDropout,
+        keras.layers.noise.GaussianDropout,
+        keras.layers.noise.GaussianNoise,
+        keras.layers.normalization.BatchNormalization,
+        keras.layers.pooling.AveragePooling1D,
+        keras.layers.pooling.AveragePooling2D,
+        keras.layers.pooling.AveragePooling3D,
+        keras.layers.pooling.GlobalAveragePooling1D,
+        keras.layers.pooling.GlobalAveragePooling2D,
+        keras.layers.pooling.GlobalAveragePooling3D,
+        keras.layers.pooling.GlobalMaxPooling1D,
+        keras.layers.pooling.GlobalMaxPooling2D,
+        keras.layers.pooling.GlobalMaxPooling3D,
+        keras.layers.pooling.MaxPooling1D,
+        keras.layers.pooling.MaxPooling2D,
+        keras.layers.pooling.MaxPooling3D,
+    )
+    return isinstance(layer, CONVNET_LAYERS)
+
+
+def is_relu_convnet_layer(layer):
+    """Checks if layer is from a convolutional network with ReLUs."""
+    return (is_convnet_layer(layer) and only_relu_activation(layer))
+
+
+def is_average_pooling(layer):
+    """Checks if layer is an average-pooling layer."""
+    AVERAGEPOOLING_LAYERS = (
+        keras.layers.pooling.AveragePooling1D,
+        keras.layers.pooling.AveragePooling2D,
+        keras.layers.pooling.AveragePooling3D,
+        keras.layers.pooling.GlobalAveragePooling1D,
+        keras.layers.pooling.GlobalAveragePooling2D,
+        keras.layers.pooling.GlobalAveragePooling3D,
+    )
+    return isinstance(layer, AVERAGEPOOLING_LAYERS)
+
+
+def is_max_pooling(layer):
+    """Checks if layer is a max-pooling layer."""
+    MAXPOOLING_LAYERS = (
+        keras.layers.pooling.MaxPooling1D,
+        keras.layers.pooling.MaxPooling2D,
+        keras.layers.pooling.MaxPooling3D,
+        keras.layers.pooling.GlobalMaxPooling1D,
+        keras.layers.pooling.GlobalMaxPooling2D,
+        keras.layers.pooling.GlobalMaxPooling3D,
+    )
+    return isinstance(layer, MAXPOOLING_LAYERS)
+
+
+def is_input_layer(layer, ignore_reshape_layers=True):
+    """Checks if layer is an input layer."""
+    # Triggers if ALL inputs of layer are connected
+    # to a Keras input layer object.
+    # Note: In the sequential api the Sequential object
+    # adds the Input layer if the user does not.
+    kgraph = get_kgraph()
+
+    layer_inputs = kgraph.get_input_layers(layer)
+    # We ignore certain layers, that do not modify
+    # the data content.
+    # todo: update this list!
+    IGNORED_LAYERS = (
+        keras.layers.Flatten,
+        keras.layers.Permute,
+        keras.layers.Reshape,
+    )
+    while any([isinstance(x, IGNORED_LAYERS) for x in layer_inputs]):
+        tmp = set()
+        for l in layer_inputs:
+            if(ignore_reshape_layers and
+               isinstance(l, IGNORED_LAYERS)):
+                tmp.update(kgraph.get_input_layers(l))
+            else:
+                tmp.add(l)
+        layer_inputs = tmp
+
+    if all([isinstance(x, keras.layers.InputLayer)
+            for x in layer_inputs]):
+        return True
+    else:
+        return False
+
+def is_layer_at_idx(layer, index, ignore_reshape_layers=True):
+    """Checks if layer is a layer at index index, by repeatedly applying is_input_layer()."""
+    kgraph = get_kgraph()

+ 1152 - 0
original_model/innvestigate/utils/keras/graph.py

@@ -0,0 +1,1152 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import range, zip
+import six
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import inspect
+import keras.backend as K
+import keras.engine.topology
+import keras.layers
+import keras.models
+import numpy as np
+
+
+from . import checks as kchecks
+from ... import layers as ilayers
+from ... import utils as iutils
+
+
+__all__ = [
+    "get_kernel",
+
+    "get_layer_inbound_count",
+    "get_layer_outbound_count",
+    "get_layer_neuronwise_io",
+    "copy_layer_wo_activation",
+    "copy_layer",
+    "pre_softmax_tensors",
+    "model_wo_softmax",
+
+    "get_model_layers",
+    "model_contains",
+
+    "trace_model_execution",
+    "get_model_execution_trace",
+    "get_model_execution_graph",
+    "print_model_execution_graph",
+
+    "get_bottleneck_nodes",
+    "get_bottleneck_tensors",
+
+    "ReverseMappingBase",
+    "reverse_model",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def get_kernel(layer):
+    """Returns the kernel weights of a layer, i.e, w/o biases."""
+    ret = [x for x in layer.get_weights() if len(x.shape) > 1]
+    assert len(ret) == 1
+    return ret[0]
+
+
+def get_input_layers(layer):
+    """Returns all layers that created this layer's inputs."""
+    ret = set()
+
+    for node_index in range(len(layer._inbound_nodes)):
+        Xs = iutils.to_list(layer.get_input_at(node_index))
+        for X in Xs:
+            ret.add(X._keras_history[0])
+
+    return ret
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def get_layer_inbound_count(layer):
+    """Returns the number inbound nodes of a layer."""
+    return len(layer._inbound_nodes)
+
+
+def get_layer_outbound_count(layer):
+    """Returns the number outbound nodes of a layer."""
+    return len(layer.outbound_nodes)
+
+
+def get_layer_neuronwise_io(layer,
+                            node_index=0,
+                            Xs=None,
+                            Ys=None,
+                            return_i=True,
+                            return_o=True):
+    """Returns the input and output for each neuron in a layer
+
+    Returns the symbolic input and output for each neuron in a layer.
+    For a dense layer this is the input output itself.
+    For convolutional layers this method extracts for each neuron
+    the input output mapping.
+
+    At the moment this function is designed
+    to work with dense and conv2d layers.
+
+    :param layer: The targeted layer.
+    :param node_index: Index of the layer node to use.
+    :param Xs: Ignore the layer's input but use Xs instead.
+    :param Ys: Ignore the layer's output but use Ys instead.
+    :param return_i: Return the inputs.
+    :param return_o: Return the outputs.
+    :return: Inputs and outputs, if specified, for each individual neuron.
+    """
+    if not kchecks.contains_kernel(layer):
+        raise NotImplementedError()
+
+    if Xs is None:
+        Xs = iutils.to_list(layer.get_input_at(node_index))
+    if Ys is None:
+        Ys = iutils.to_list(layer.get_output_at(node_index))
+
+    if isinstance(layer, keras.layers.Dense):
+        # Xs and Ys are already in shape.
+        ret_Xs = Xs
+        ret_Ys = Ys
+    elif isinstance(layer, keras.layers.Conv2D):
+        kernel = get_kernel(layer)
+        # Expect filter dimension to be last.
+        n_channels = kernel.shape[-1]
+
+        if return_i:
+            extract_patches = ilayers.ExtractConv2DPatches(kernel.shape[:2],
+                                                           kernel.shape[2],
+                                                           layer.strides,
+                                                           layer.dilation_rate,
+                                                           layer.padding)
+            # shape [samples, out_row, out_col, weight_size]
+            reshape = ilayers.Reshape((-1, np.product(kernel.shape[:3])))
+            ret_Xs = [reshape(extract_patches(x)) for x in Xs]
+
+        if return_o:
+            # Get Ys into shape (samples, channels)
+            if K.image_data_format() == "channels_first":
+                # Ys shape is [samples, channels, out_row, out_col]
+                def reshape(x):
+                    x = ilayers.Transpose((0, 2, 3, 1))(x)
+                    x = ilayers.Reshape((-1, n_channels))(x)
+                    return x
+            else:
+                # Ys shape is [samples, out_row, out_col, channels]
+                def reshape(x):
+                    x = ilayers.Reshape((-1, n_channels))(x)
+                    return x
+            ret_Ys = [reshape(x) for x in Ys]
+
+    else:
+        raise NotImplementedError()
+
+    # Xs is (n, d) and Ys is (d, channels)
+    if return_i and return_o:
+        return ret_Xs, ret_Ys
+    elif return_i:
+        return ret_Xs
+    elif return_o:
+        return ret_Ys
+    else:
+        raise Exception()
+
+
+def get_symbolic_weight_names(layer, weights=None):
+    """Attribute names for weights
+
+    Looks up the attribute names of weight tensors.
+
+    :param layer: Targeted layer.
+    :param weights: A list of weight tensors.
+    :return: The attribute names of the weights.
+    """
+
+    if weights is None:
+        weights = layer.weights
+
+    good_guesses = [
+        "kernel",
+        "bias",
+        "gamma",
+        "beta",
+        "moving_mean",
+        "moving_variance",
+        "depthwise_kernel",
+        "pointwise_kernel"
+    ]
+
+    ret = []
+    for weight in weights:
+        for attr_name in good_guesses+dir(layer):
+            if(hasattr(layer, attr_name) and
+               id(weight) == id(getattr(layer, attr_name))):
+                ret.append(attr_name)
+                break
+    if len(weights) != len(ret):
+        raise Exception("Could not find symoblic weight name(s).")
+
+    return ret
+
+
+def update_symbolic_weights(layer, weight_mapping):
+    """Updates the symbolic tensors of a layer
+
+    Updates the symbolic tensors of a layer by replacing them.
+
+    Note this does not update the loss or anything alike!
+    Use with caution!
+
+    :param layer: Targeted layer.
+    :param weight_mapping: Dict with attribute name and weight tensors
+      as keys and values.
+    """
+
+    trainable_weight_ids = [id(x) for x in layer._trainable_weights]
+    non_trainable_weight_ids = [id(x) for x in layer._non_trainable_weights]
+
+    for name, weight in six.iteritems(weight_mapping):
+        current_weight = getattr(layer, name)
+        current_weight_id = id(current_weight)
+
+        if current_weight_id in trainable_weight_ids:
+            idx = trainable_weight_ids.index(current_weight_id)
+            layer._trainable_weights[idx] = weight
+        else:
+            idx = non_trainable_weight_ids.index(current_weight_id)
+            layer._non_trainable_weights[idx] = weight
+
+        setattr(layer, name, weight)
+
+
+def get_layer_from_config(old_layer,
+                          new_config,
+                          weights=None,
+                          reuse_symbolic_tensors=True):
+    """Creates a new layer from a config
+
+    Creates a new layer given a changed config and weights etc.
+
+    :param old_layer: A layer that shall be used as base.
+    :param new_config: The config to create the new layer.
+    :param weights: Weights to set in the new layer.
+      Options: np tensors, symbolic tensors, or None,
+      in which case the weights from old_layers are used.
+    :param reuse_symbolic_tensors: If the weights of the
+      old_layer are used copy the symbolic ones or copy
+      the Numpy weights.
+    :return: The new layer instance.
+    """
+    new_layer = old_layer.__class__.from_config(new_config)
+
+    if weights is None:
+        if reuse_symbolic_tensors:
+            weights = old_layer.weights
+        else:
+            weights = old_layer.get_weights()
+
+    if len(weights) > 0:
+        input_shapes = old_layer.get_input_shape_at(0)
+        # todo: inspect and set initializers to something fast for speedup
+        new_layer.build(input_shapes)
+
+        is_np_weight = [isinstance(x, np.ndarray) for x in weights]
+        if all(is_np_weight):
+            new_layer.set_weights(weights)
+        else:
+            if any(is_np_weight):
+                raise ValueError("Expect either all weights to be "
+                                 "np tensors or symbolic tensors.")
+
+            symbolic_names = get_symbolic_weight_names(old_layer)
+            update = {name: weight
+                      for name, weight in zip(symbolic_names, weights)}
+            update_symbolic_weights(new_layer, update)
+
+    return new_layer
+
+
+def copy_layer_wo_activation(layer,
+                             keep_bias=True,
+                             name_template=None,
+                             weights=None,
+                             reuse_symbolic_tensors=True,
+                             **kwargs):
+    """Copy a Keras layer and remove the activations
+
+    Copies a Keras layer but remove potential activations.
+
+    :param layer: A layer that should be copied.
+    :param keep_bias: Keep a potential bias.
+    :param weights: Weights to set in the new layer.
+      Options: np tensors, symbolic tensors, or None,
+      in which case the weights from old_layers are used.
+    :param reuse_symbolic_tensors: If the weights of the
+      old_layer are used copy the symbolic ones or copy
+      the Numpy weights.
+    :return: The new layer instance.
+    """
+    config = layer.get_config()
+    if name_template is None:
+        config["name"] = None
+    else:
+        config["name"] = name_template % config["name"]
+    if kchecks.contains_activation(layer):
+        config["activation"] = None
+    if hasattr(layer, "use_bias"):
+        if keep_bias is False and config.get("use_bias", True):
+            config["use_bias"] = False
+            if weights is None:
+                if reuse_symbolic_tensors:
+                    weights = layer.weights[:-1]
+                else:
+                    weights = layer.get_weights()[:-1]
+    return get_layer_from_config(layer, config, weights=weights, **kwargs)
+
+
+def copy_layer(layer,
+               keep_bias=True,
+               name_template=None,
+               weights=None,
+               reuse_symbolic_tensors=True,
+               **kwargs):
+    """Copy a Keras layer
+
+    Copies a Keras layer.
+
+    :param layer: A layer that should be copied.
+    :param keep_bias: Keep a potential bias.
+    :param weights: Weights to set in the new layer.
+      Options: np tensors, symbolic tensors, or None,
+      in which case the weights from old_layers are used.
+    :param reuse_symbolic_tensors: If the weights of the
+      old_layer are used copy the symbolic ones or copy
+      the Numpy weights.
+    :return: The new layer instance.
+    """
+    config = layer.get_config()
+    if name_template is None:
+        config["name"] = None
+    else:
+        config["name"] = name_template % config["name"]
+    if hasattr(layer, "use_bias"):
+        if keep_bias is False and config.get("use_bias", True):
+            config["use_bias"] = False
+            if weights is None:
+                if reuse_symbolic_tensors:
+                    weights = layer.weights[:-1]
+                else:
+                    weights = layer.get_weights()[:-1]
+    return get_layer_from_config(layer, config, weights=weights, **kwargs)
+
+
+def pre_softmax_tensors(Xs, should_find_softmax=True):
+    """Finds the tensors that were preceeding a potential softmax."""
+    softmax_found = False
+
+    Xs = iutils.to_list(Xs)
+    ret = []
+    for x in Xs:
+        layer, node_index, tensor_index = x._keras_history
+        if kchecks.contains_activation(layer, activation="softmax"):
+            softmax_found = True
+            if isinstance(layer, keras.layers.Activation):
+                ret.append(layer.get_input_at(node_index))
+            else:
+                layer_wo_act = copy_layer_wo_activation(layer)
+                ret.append(layer_wo_act(layer.get_input_at(node_index)))
+
+    if should_find_softmax and not softmax_found:
+        raise Exception("No softmax found.")
+
+    return ret
+
+
+def model_wo_softmax(model):
+    """Creates a new model w/o the final softmax activation."""
+    return keras.models.Model(inputs=model.inputs,
+                              outputs=pre_softmax_tensors(model.outputs),
+                              name=model.name)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def get_model_layers(model):
+    """Returns all layers of a model."""
+    ret = []
+
+    def collect_layers(container):
+        for layer in container.layers:
+            assert layer not in ret
+            ret.append(layer)
+            if kchecks.is_network(layer):
+                collect_layers(layer)
+    collect_layers(model)
+
+    return ret
+
+
+def model_contains(model, layer_condition, return_only_counts=False):
+    if callable(layer_condition):
+        layer_condition = [layer_condition, ]
+        single_condition = True
+    else:
+        single_condition = False
+
+    layers = get_model_layers(model)
+    collected_layers = []
+    for condition in layer_condition:
+        tmp = [layer for layer in layers if condition(layer)]
+        collected_layers.append(tmp)
+    if return_only_counts is True:
+        collected_layers = [len(v) for v in collected_layers]
+
+    if single_condition is True:
+        return collected_layers[0]
+    else:
+        return collected_layers
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def apply_mapping_to_fused_bn_layer(mapping, fuse_mode="one_linear"):
+    """
+    Applies a mapping to a linearized Batch Normalization layer.
+
+    :param mapping: The mapping to be applied.
+      Should take parameters layer and reverse_state and
+      return a mapping function.
+    :param fuse_mode: Either 'one_linear': apply the mapping
+      to a once linearized layer, or
+      'two_linear': apply to twice to a twice linearized layer.
+    """
+    if fuse_mode not in ["one_linear", "two_linear"]:
+        raise ValueError("fuse_mode can only be 'one_linear' or 'two_linear'")
+
+    # todo(alber): remove this workaround and make a proper class
+    def ScaleLayer(kernel, bias):
+        _kernel = kernel
+        _bias = bias
+
+        class ScaleLayer(keras.layers.Layer):
+
+            def __init__(self, use_bias=True, **kwargs):
+                self._kernel_to_be = _kernel
+                self._bias_to_be = _bias
+                self.use_bias = use_bias
+                super(ScaleLayer, self).__init__(**kwargs)
+
+            def build(self, input_shape):
+                self.kernel = self.add_weight(
+                    name='kernel',
+                    shape=K.int_shape(self._kernel_to_be),
+                    initializer=lambda a, b=None: self._kernel_to_be,
+                    trainable=False)
+                if self.use_bias:
+                    self.bias = self.add_weight(
+                        name='bias',
+                        shape=K.int_shape(self._bias_to_be),
+                        initializer=lambda a, b=None: self._bias_to_be,
+                        trainable=False)
+                super(ScaleLayer, self).build(input_shape)
+
+            def call(self, x):
+                ret = (x * self.kernel)
+                if self.use_bias:
+                    ret += self.bias
+                return ret
+
+            def compute_output_shape(self, input_shape):
+                return input_shape
+
+        return ScaleLayer()
+
+    def meta_mapping(layer, reverse_state):
+        # get bn params
+        weights = layer.weights[:]
+        if layer.scale:
+            gamma = weights.pop(0)
+        else:
+            gamma = K.ones_like(weights[0])
+        if layer.center:
+            beta = weights.pop(0)
+        else:
+            beta = K.zeros_like(weights[0])
+        mean, variance = weights
+
+        if fuse_mode == "one_linear":
+            tmp = K.sqrt(variance**2 + layer.epsilon)
+            tmp_k = gamma / tmp
+            tmp_b = -mean / tmp + beta
+
+            inputs = layer.get_input_at(0)
+            surrogate_layer = ScaleLayer(tmp_k, tmp_b)
+            # init layer
+            surrogate_layer(inputs)
+            actual_mapping = mapping(surrogate_layer, reverse_state).apply
+        else:
+            tmp = K.sqrt(variance**2 + layer.epsilon)
+            tmp_k1 = 1 / tmp
+            tmp_b1 = -mean / tmp
+            tmp_k2 = gamma
+            tmp_b2 = beta
+
+            inputs = layer.get_input_at(0)
+            surrogate_layer1 = ScaleLayer(tmp_k1, tmp_b1)
+            surrogate_layer2 = ScaleLayer(tmp_k2, tmp_b2)
+            # init layers
+            surrogate_layer1(inputs)
+            surrogate_layer2(inputs)
+            # todo(alber): update reverse state
+            actual_mapping_1 = mapping(surrogate_layer1, reverse_state).apply
+            actual_mapping_2 = mapping(surrogate_layer2, reverse_state).apply
+
+            def actual_mapping(Xs, Ys, reversed_Ys, reverse_state):
+                from . import apply as kapply
+                X2s = kapply(surrogate_layer1, Xs)
+                # Apply first mapping
+                # todo(alber): update reverse state
+                reversed_X2s = actual_mapping_2(
+                    X2s, Ys, reversed_Ys, reverse_state)
+                return actual_mapping_1(Xs, X2s, reversed_X2s, reverse_state)
+        return actual_mapping
+    return meta_mapping
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def trace_model_execution(model, reapply_on_copied_layers=False):
+    """
+    Trace and linearize excecution of a model and it's possible containers.
+    Return a triple with all layers, a list with a linearized execution
+    with (layer, input_tensors, output_tensors), and, possible regenerated,
+    outputs of the exectution.
+
+    :param model: A kera model.
+    :param reapply_on_copied_layers: If the execution needs to be linearized,
+      reapply with copied layers. Might be slow. Prevents changes of the
+      original layer's node lists.
+    """
+
+    # Get all layers in model.
+    layers = get_model_layers(model)
+
+    # Check if some layers are containers.
+    # Ignoring the outermost container, i.e. the passed model.
+    contains_container = any([((l is not model) and kchecks.is_network(l))
+                              for l in layers])
+
+    # If so rebuild the graph, otherwise recycle computations,
+    # and create executed node list. (Keep track of paths?)
+    if contains_container is True:
+        # When containers/models are used as layers, then layers
+        # inside the container/model do not keep track of nodes.
+        # This makes it impossible to iterate of the nodes list and
+        # recover the input output tensors. (see else clause)
+        #
+        # To recover the computational graph we need to re-apply it.
+        # This implies that the tensors-object we use for the forward
+        # pass are different to the passed model. This it not the case
+        # for the else clause.
+        #
+        # Note that reapplying the model does only change the inbound
+        # and outbound nodes of the model itself. We copy the model
+        # so the passed model should not be affected from the
+        # reapplication.
+        executed_nodes = []
+
+        # Monkeypatch the call function in all the used layer classes.
+        monkey_patches = [(layer, getattr(layer, "call")) for layer in layers]
+        try:
+            def patch(self, method):
+                if hasattr(method, "__patched__") is True:
+                    raise Exception("Should not happen as we patch "
+                                    "objects not classes.")
+
+                def f(*args, **kwargs):
+                    input_tensors = args[0]
+                    output_tensors = method(*args, **kwargs)
+                    executed_nodes.append((self,
+                                           input_tensors,
+                                           output_tensors))
+                    return output_tensors
+                f.__patched__ = True
+                return f
+
+            # Apply the patches.
+            for layer in layers:
+                setattr(layer, "call", patch(layer, getattr(layer, "call")))
+
+            # Trigger reapplication of model.
+            model_copy = keras.models.Model(inputs=model.inputs,
+                                            outputs=model.outputs)
+            outputs = iutils.to_list(model_copy(model.inputs))
+        finally:
+            # Revert the monkey patches
+            for layer, old_method in monkey_patches:
+                setattr(layer, "call", old_method)
+
+        # Now we have the problem that all the tensors
+        # do not have a keras_history attribute as they are not part
+        # of any node. Apply the flat model to get it.
+        from . import apply as kapply
+        new_executed_nodes = []
+        tensor_mapping = {tmp: tmp for tmp in model.inputs}
+        if reapply_on_copied_layers is True:
+            layer_mapping = {layer: copy_layer(layer) for layer in layers}
+        else:
+            layer_mapping = {layer: layer for layer in layers}
+
+        for layer, Xs, Ys in executed_nodes:
+            layer = layer_mapping[layer]
+            Xs, Ys = iutils.to_list(Xs), iutils.to_list(Ys)
+
+            if isinstance(layer, keras.layers.InputLayer):
+                # Special case. Do nothing.
+                new_Xs, new_Ys = Xs, Ys
+            else:
+                new_Xs = [tensor_mapping[x] for x in Xs]
+                new_Ys = iutils.to_list(kapply(layer, new_Xs))
+
+            tensor_mapping.update({k: v for k, v in zip(Ys, new_Ys)})
+            new_executed_nodes.append((layer, new_Xs, new_Ys))
+
+        layers = [layer_mapping[layer] for layer in layers]
+        outputs = [tensor_mapping[x] for x in outputs]
+        executed_nodes = new_executed_nodes
+    else:
+        # Easy and safe way.
+        reverse_executed_nodes = [
+            (node.outbound_layer, node.input_tensors, node.output_tensors)
+            for depth in sorted(model._nodes_by_depth.keys())
+            for node in model._nodes_by_depth[depth]
+        ]
+        outputs = model.outputs
+
+        executed_nodes = reversed(reverse_executed_nodes)
+
+    # This list contains potentially nodes that are not part
+    # final execution graph.
+    # E.g., a layer was also applied outside of the model. Then its
+    # node list contains nodes that do not contribute to the model's output.
+    # Those nodes are filtered here.
+    used_as_input = [x for x in outputs]
+    tmp = []
+    for l, Xs, Ys in reversed(list(executed_nodes)):
+        if all([y in used_as_input for y in Ys]):
+            used_as_input += Xs
+            tmp.append((l, Xs, Ys))
+    executed_nodes = list(reversed(tmp))
+
+    return layers, executed_nodes, outputs
+
+
+def get_model_execution_trace(model,
+                              keep_input_layers=False,
+                              reapply_on_copied_layers=False):
+    """
+    Returns a list representing the execution graph.
+    Each key is the node's id as it is used by the reverse_model method.
+
+    Each associated value contains a dictionary with the following items:
+
+    * nid: the node id.
+    * layer: the layer creating this node.
+    * Xs: the input tensors (only valid if not in a nested container).
+    * Ys: the output tensors (only valid if not in a nested container).
+    * Xs_nids: the ids of the nodes creating the Xs.
+    * Ys_nids: the ids of nodes using the according output tensor.
+    * Xs_layers: the layer that created the accodring input tensor.
+    * Ys_layers: the layers using the according output tensor.
+
+    :param model: A kera model.
+    :param keep_input_layers: Keep input layers.
+    :param reapply_on_copied_layers: If the execution needs to be linearized,
+      reapply with copied layers. Might be slow. Prevents changes of the
+      original layer's node lists.
+    """
+    _, execution_trace, _ = trace_model_execution(
+        model,
+        reapply_on_copied_layers=reapply_on_copied_layers)
+
+    # Enrich trace with node ids.
+    current_nid = 0
+    tmp = []
+    for l, Xs, Ys in execution_trace:
+        if isinstance(l, keras.layers.InputLayer):
+            tmp.append((None, l, Xs, Ys))
+        else:
+            tmp.append((current_nid, l, Xs, Ys))
+            current_nid += 1
+    execution_trace = tmp
+
+    # Create lookups from tensor to creating or receiving layer-node
+    inputs_to_node = {}
+    outputs_to_node = {}
+    for nid, l, Xs, Ys in execution_trace:
+        if nid is not None:
+            for X in Xs:
+                Xid = id(X)
+                if Xid in inputs_to_node:
+                    inputs_to_node[Xid].append(nid)
+                else:
+                    inputs_to_node[Xid] = [nid]
+
+        if keep_input_layers or nid is not None:
+            for Y in Ys:
+                Yid = id(Y)
+                if Yid in inputs_to_node:
+                    raise Exception("Cannot be more than one creating node.")
+                outputs_to_node[Yid] = nid
+
+    # Enrich trace with this info.
+    nid_to_nodes = {t[0]: t for t in execution_trace}
+    tmp = []
+    for nid, l, Xs, Ys in execution_trace:
+        if isinstance(l, keras.layers.InputLayer):
+            # The nids that created or receive the tensors.
+            Xs_nids = []  # Input layer does not receive.
+            Ys_nids = [inputs_to_node[id(Y)] for Y in Ys]
+            # The layers that created or receive the tensors.
+            Xs_layers = []  # Input layer does not receive.
+            Ys_layers = [[nid_to_nodes[Ynid][1] for Ynid in Ynids2]
+                         for Ynids2 in Ys_nids]
+        else:
+            # The nids that created or receive the tensors.
+            Xs_nids = [outputs_to_node.get(id(X), None) for X in Xs]
+            Ys_nids = [inputs_to_node.get(id(Y), [None]) for Y in Ys]
+            # The layers that created or receive the tensors.
+            Xs_layers = [nid_to_nodes[Xnid][1]
+                         for Xnid in Xs_nids if Xnid is not None]
+            Ys_layers = [[nid_to_nodes[Ynid][1]
+                          for Ynid in Ynids2 if Ynid is not None]
+                         for Ynids2 in Ys_nids]
+
+        entry = {
+            "nid": nid,
+            "layer": l,
+            "Xs": Xs,
+            "Ys": Ys,
+            "Xs_nids": Xs_nids,
+            "Ys_nids": Ys_nids,
+            "Xs_layers": Xs_layers,
+            "Ys_layers": Ys_layers,
+        }
+        tmp.append(entry)
+    execution_trace = tmp
+
+    if not keep_input_layers:
+        execution_trace = [tmp
+                           for tmp in execution_trace
+                           if tmp["nid"] is not None]
+
+    return execution_trace
+
+
+def get_model_execution_graph(model, keep_input_layers=False):
+    """
+    Returns a dictionary representing the execution graph.
+    Each key is the node's id as it is used by the reverse_model method.
+
+    Each associated value contains a dictionary with the following items:
+
+    * nid: the node id.
+    * layer: the layer creating this node.
+    * Xs: the input tensors (only valid if not in a nested container).
+    * Ys: the output tensors (only valid if not in a nested container).
+    * Xs_nids: the ids of the nodes creating the Xs.
+    * Ys_nids: the ids of nodes using the according output tensor.
+    * Xs_layers: the layer that created the accodring input tensor.
+    * Ys_layers: the layers using the according output tensor.
+
+    :param model: A kera model.
+    :param keep_input_layers: Keep input layers.
+    """
+    trace = get_model_execution_trace(model,
+                                      keep_input_layers=keep_input_layers,
+                                      reapply_on_copied_layers=False)
+
+    input_layers = [tmp for tmp in trace if tmp["nid"] is None]
+    graph = {tmp["nid"]: tmp for tmp in trace}
+    if keep_input_layers:
+        graph[None] = input_layers
+
+    return graph
+
+
+def print_model_execution_graph(graph):
+    """Pretty print of a model execution graph."""
+
+    def nids_as_str(nids):
+        return ", ".join(["%s" % nid for nid in nids])
+
+    def print_node(node):
+        print("  [NID: %4s] [Layer: %20s] "
+              "[Inputs from: %20s] [Outputs to: %20s]" %
+              (node["nid"],
+               node["layer"].name,
+               nids_as_str(node["Xs_nids"]),
+               nids_as_str(node["Ys_nids"]),))
+
+    if None in graph:
+        print("Graph input layers:")
+        for tmp in graph[None]:
+            print_node(tmp)
+
+    print("Graph nodes:")
+    for nid in sorted([k for k in graph if k is not None]):
+        if nid is None:
+            continue
+        print_node(graph[nid])
+
+
+def get_bottleneck_nodes(inputs, outputs, execution_list):
+    """
+    Given an execution list this function returns all nodes that
+    are a bottleneck in the network, i.e., "all information" must pass
+    through this node.
+    """
+
+    forward_connections = {}
+    for l, Xs, Ys in execution_list:
+        if isinstance(l, keras.layers.InputLayer):
+            # Special case, do nothing.
+            continue
+
+        for x in Xs:
+            if x in forward_connections:
+                forward_connections[x] += Ys
+            else:
+                forward_connections[x] = list(Ys)
+
+    open_connections = {}
+    for x in inputs:
+        for fw_c in forward_connections[x]:
+            open_connections[fw_c] = True
+
+    ret = list()
+    for l, Xs, Ys in execution_list:
+        if isinstance(l, keras.layers.InputLayer):
+            # Special case, do nothing.
+            # Note: if a single input branches
+            # this is not detected.
+            continue
+
+        for y in Ys:
+            assert y in open_connections
+            del open_connections[y]
+
+        if len(open_connections) == 0:
+            ret.append((l, (Xs, Ys)))
+
+        for y in Ys:
+            if y not in outputs:
+                for fw_c in forward_connections[y]:
+                    open_connections[fw_c] = True
+
+    return ret
+
+
+def get_bottleneck_tensors(inputs, outputs, execution_list):
+    """
+    Given an execution list this function returns all tensors that
+    are a bottleneck in the network, i.e., "all information" must pass
+    through this tensor.
+    """
+
+    nodes = get_bottleneck_nodes(inputs, outputs, execution_list)
+
+    ret = list()
+    for l, (Xs, Ys) in nodes:
+        for tensor_list in (Xs, Ys):
+            if len(tensor_list) == 1:
+                tensor = tensor_list[0]
+                if tensor not in ret:
+                    ret.append(tensor)
+            else:
+                # TODO(albermax): put warning here?
+                pass
+    return ret
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class ReverseMappingBase(object):
+
+    def __init__(self, layer, state):
+        pass
+
+    def apply(self, Xs, Yx, reversed_Ys, reverse_state):
+        raise NotImplementedError()
+
+
+def reverse_model(model, reverse_mappings,
+                  default_reverse_mapping=None,
+                  head_mapping=None,
+                  stop_mapping_at_tensors=[],
+                  verbose=False,
+                  return_all_reversed_tensors=False,
+                  clip_all_reversed_tensors=False,
+                  project_bottleneck_tensors=False,
+                  execution_trace=None,
+                  reapply_on_copied_layers=False):
+    """
+    Reverses a Keras model based on the given reverse functions.
+    It returns the reverted tensors for the according model inputs.
+
+    :param model: A Keras model.
+    :param reverse_mappings: Either a callable that matches layers to
+      mappings or a dictionary with layers as keys and mappings as values.
+      Allowed as mapping forms are:
+          * A function of form (A) f(Xs, Ys, reversed_Ys, reverse_state).
+          * A function of form f(B) f(layer, reverse_state) that returns
+            a function of form (A).
+          * A :class:`ReverseMappingBase` subclass.
+    :param default_reverse_mapping: A function that reverses layers for
+      which no mapping was given by param "reverse_mappings".
+    :param head_mapping: Map output tensors to new values before passing
+      them into the reverted network.
+    :param stop_mapping_at_tensors: Tensors at which to stop the mapping.
+      Similar to stop_gradient parameters for gradient computation.
+    :param verbose: Print what's going on.
+    :param return_all_reversed_tensors: Return all reverted tensors in addition
+      to reverted model input tensors.
+    :param clip_all_reversed_tensors: Clip each reverted tensor. False or tuple
+      with min/max value.
+    :param project_bottleneck_tensors: Project bottleneck layers in the
+      reverting process into a given value range. False, True or (a, b) for
+      projection range.
+    :param reapply_on_copied_layers: When a model execution needs to
+      linearized and copy layers before reapplying them. See
+      :func:`trace_model_execution`.
+    """
+
+    # Set default values ######################################################
+
+    if head_mapping is None:
+        def head_mapping(X):
+            return X
+
+    if not callable(reverse_mappings):
+        # not callable, assume a dict that maps from layer to mapping
+        reverse_mapping_data = reverse_mappings
+
+        def reverse_mappings(layer):
+            try:
+                return reverse_mapping_data[type(layer)]
+            except KeyError:
+                return None
+
+    def _print(s):
+        if verbose is True:
+            print(s)
+
+    # Initialize structure that keeps track of reversed tensors ###############
+
+    reversed_tensors = {}
+    bottleneck_tensors = set()
+
+    def add_reversed_tensors(nid,
+                             tensors_list,
+                             reversed_tensors_list):
+
+        def add_reversed_tensor(i, X, reversed_X):
+            # Do not keep tensors that should stop the mapping.
+            if X in stop_mapping_at_tensors:
+                return
+
+            if X not in reversed_tensors:
+                reversed_tensors[X] = {"id": (nid, i),
+                                       "tensor": reversed_X}
+            else:
+                tmp = reversed_tensors[X]
+                if "tensor" in tmp and "tensors" in tmp:
+                    raise Exception("Wrong order, tensors already aggregated!")
+                if "tensor" in tmp:
+                    tmp["tensors"] = [tmp["tensor"], reversed_X]
+                    del tmp["tensor"]
+                else:
+                    tmp["tensors"].append(reversed_X)
+
+        tmp = zip(tensors_list, reversed_tensors_list)
+        for i, (X, reversed_X) in enumerate(tmp):
+            add_reversed_tensor(i, X, reversed_X)
+
+    def get_reversed_tensor(tensor):
+        tmp = reversed_tensors[tensor]
+
+        if "final_tensor" not in tmp:
+            if "tensor" not in tmp:
+                final_tensor = keras.layers.Add()(tmp["tensors"])
+            else:
+                final_tensor = tmp["tensor"]
+
+            if project_bottleneck_tensors is not False:
+                if tensor in bottleneck_tensors:
+                    project = ilayers.Project(project_bottleneck_tensors)
+                    final_tensor = project(final_tensor)
+
+            if clip_all_reversed_tensors is not False:
+                clip = ilayers.Clip(*clip_all_reversed_tensors)
+                final_tensor = clip(final_tensor)
+
+            tmp["final_tensor"] = final_tensor
+
+        return tmp["final_tensor"]
+
+    # Reverse the model #######################################################
+    _print("Reverse model: {}".format(model))
+
+    # Create a list with nodes in reverse execution order.
+    if execution_trace is None:
+        execution_trace = trace_model_execution(
+            model,
+            reapply_on_copied_layers=reapply_on_copied_layers)
+    layers, execution_list, outputs = execution_trace
+    len_execution_list = len(execution_list)
+    num_input_layers = len([_ for l, _, _ in execution_list
+                            if isinstance(l, keras.layers.InputLayer)])
+    len_execution_list_wo_inputs_layers = len_execution_list - num_input_layers
+    reverse_execution_list = reversed(execution_list)
+
+    # Initialize the reverse mapping functions.
+    initialized_reverse_mappings = {}
+    for layer in layers:
+        # A layer can be shared, i.e., applied several times.
+        # Allow to share a ReverMappingBase for each layer instance
+        # in order to reduce the overhead.
+
+        meta_reverse_mapping = reverse_mappings(layer)
+        if meta_reverse_mapping is None:
+            reverse_mapping = default_reverse_mapping
+        elif(inspect.isclass(meta_reverse_mapping) and
+             issubclass(meta_reverse_mapping, ReverseMappingBase)):
+            # Mapping is a class
+            reverse_mapping_obj = meta_reverse_mapping(
+                layer,
+                {
+                    "model": model,
+                    "layer": layer,
+                }
+            )
+            reverse_mapping = reverse_mapping_obj.apply
+        else:
+            def parameter_count(func):
+                if hasattr(inspect, "signature"):
+                    ret = len(inspect.signature(func).parameters)
+                else:
+                    spec = inspect.getargspec(func)
+                    ret = len(spec.args)
+                    if spec.varargs is not None:
+                        ret += len(spec.varargs)
+                    if spec.keywords is not None:
+                        ret += len(spec.keywords)
+                    if ret == 3:
+                        # assume class function with self
+                        ret -= 1
+                return ret
+
+            if(callable(meta_reverse_mapping) and
+               parameter_count(meta_reverse_mapping) == 2):
+                # Function that returns mapping
+                reverse_mapping = meta_reverse_mapping(
+                    layer,
+                    {
+                        "model": model,
+                        "layer": layer,
+                    }
+                )
+            else:
+                # Nothing meta here
+                reverse_mapping = meta_reverse_mapping
+
+        initialized_reverse_mappings[layer] = reverse_mapping
+
+    if project_bottleneck_tensors:
+        bottleneck_tensors.update(
+            get_bottleneck_tensors(
+                model.inputs,
+                outputs,
+                execution_list))
+
+    # Initialize the reverse tensor mappings.
+    add_reversed_tensors(-1,
+                         outputs,
+                         [head_mapping(tmp) for tmp in outputs])
+
+    # Follow the list and revert the graph.
+    for _nid, (layer, Xs, Ys) in enumerate(reverse_execution_list):
+        nid = len_execution_list_wo_inputs_layers - _nid - 1
+
+        if isinstance(layer, keras.layers.InputLayer):
+            # Special case. Do nothing.
+            pass
+        elif kchecks.is_network(layer):
+            raise Exception("This is not supposed to happen!")
+        else:
+            Xs, Ys = iutils.to_list(Xs), iutils.to_list(Ys)
+            if not all([ys in reversed_tensors for ys in Ys]):
+                # This node is not part of our computational graph.
+                # The (node-)world is bigger than this model.
+                # Potentially this node is also not part of the
+                # reversed tensor set because it depends on a tensor
+                # that is listed in stop_mapping_at_tensors.
+                continue
+            reversed_Ys = [get_reversed_tensor(ys)
+                           for ys in Ys]
+            local_stop_mapping_at_tensors = [x for x in Xs
+                                             if x in stop_mapping_at_tensors]
+
+            _print("  [NID: {}] Reverse layer-node {}".format(nid, layer))
+            reverse_mapping = initialized_reverse_mappings[layer]
+            reversed_Xs = reverse_mapping(
+                Xs, Ys, reversed_Ys,
+                {
+                    "nid": nid,
+                    "model": model,
+                    "layer": layer,
+                    "stop_mapping_at_tensors": local_stop_mapping_at_tensors,
+                })
+            reversed_Xs = iutils.to_list(reversed_Xs)
+            add_reversed_tensors(nid, Xs, reversed_Xs)
+
+    # Return requested values #################################################
+
+    #THIS LINE ADDED FOR 3 INPUT MODEL:
+    stop_mapping_at_tensors = [model.inputs[1],model.inputs[2]]
+    
+    reversed_input_tensors = [get_reversed_tensor(tmp)
+                              for tmp in model.inputs
+                              if tmp not in stop_mapping_at_tensors]
+    if return_all_reversed_tensors is True:
+        return reversed_input_tensors, reversed_tensors
+    else:
+        return reversed_input_tensors

+ 0 - 0
original_model/innvestigate/utils/tests/__init__.py


+ 338 - 0
original_model/innvestigate/utils/tests/dryrun.py

@@ -0,0 +1,338 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+import six
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.backend as K
+import keras.models
+import numpy as np
+import unittest
+
+from ...analyzer.base import AnalyzerBase
+from . import networks
+
+
+__all__ = [
+    "AnalyzerTestCase",
+    "EqualAnalyzerTestCase",
+    "PatternComputerTestCase",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def _set_zero_weights_to_random(weights):
+    ret = []
+    for weight in weights:
+        if weight.sum() == 0:
+            weight = np.random.rand(*weight.shape)
+        ret.append(weight)
+    return ret
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class BaseLayerTestCase(unittest.TestCase):
+    """
+    A dryrun test on various networks for an analyzing method.
+
+    For each network the test check that the generated network
+    has the right output shape, can be compiled
+    and executed with random inputs.
+    """
+
+    _network_filter = "trivia.*"
+
+    def __init__(self, *args, **kwargs):
+        network_filter = kwargs.pop("network_filter", None)
+        if network_filter is not None:
+            self._network_filter = network_filter
+        super(BaseLayerTestCase, self).__init__(*args, **kwargs)
+
+    def _apply_test(self, network):
+        raise NotImplementedError("Set in subclass.")
+
+    def runTest(self):
+        np.random.seed(2349784365)
+        K.clear_session()
+
+        for network in networks.iterator(self._network_filter,
+                                         clear_sessions=True):
+            if six.PY2:
+                self._apply_test(network)
+            else:
+                with self.subTest(network_name=network["name"]):
+                    self._apply_test(network)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class AnalyzerTestCase(BaseLayerTestCase):
+    """TestCase for analyzers execution
+
+    TestCase that applies the method to several networks and
+    runs the analyzer with random data.
+
+    :param method: A function that returns an Analyzer class.
+    """
+    def __init__(self, *args, **kwargs):
+        method = kwargs.pop("method", None)
+        if method is not None:
+            self._method = method
+        super(AnalyzerTestCase, self).__init__(*args, **kwargs)
+
+    def _method(self, model):
+        raise NotImplementedError("Set in subclass.")
+
+    def _apply_test(self, network):
+        # Create model.
+        model = keras.models.Model(inputs=network["in"],
+                                   outputs=network["out"])
+        model.set_weights(_set_zero_weights_to_random(model.get_weights()))
+        # Get analyzer.
+        analyzer = self._method(model)
+        # Dryrun.
+        x = np.random.rand(1, *(network["input_shape"][1:]))
+        analysis = analyzer.analyze(x)
+        self.assertEqual(tuple(analysis.shape),
+                         (1,)+tuple(network["input_shape"][1:]))
+        self.assertFalse(np.any(np.isinf(analysis.ravel())))
+        self.assertFalse(np.any(np.isnan(analysis.ravel())))
+
+
+def test_analyzer(method, network_filter):
+    """Workaround for move from unit-tests to pytest."""
+    # todo: Mixing of pytest and unittest is not ideal.
+    # Move completely to pytest.
+    test_case = AnalyzerTestCase(method=method,
+                                 network_filter=network_filter)
+    test_result = unittest.TextTestRunner().run(test_case)
+    assert len(test_result.errors) == 0
+    assert len(test_result.failures) == 0
+
+
+class AnalyzerTrainTestCase(BaseLayerTestCase):
+    """TestCase for analyzers execution
+
+    TestCase that applies the method to several networks and
+    trains and runs the analyzer with random data.
+
+    :param method: A function that returns an Analyzer class.
+    """
+
+    def __init__(self, *args, **kwargs):
+        method = kwargs.pop("method", None)
+        if method is not None:
+            self._method = method
+        super(AnalyzerTrainTestCase, self).__init__(*args, **kwargs)
+
+    def _method(self, model):
+        raise NotImplementedError("Set in subclass.")
+
+    def _apply_test(self, network):
+        # Create model.
+        model = keras.models.Model(inputs=network["in"],
+                                   outputs=network["out"])
+        model.set_weights(_set_zero_weights_to_random(model.get_weights()))
+        # Get analyzer.
+        analyzer = self._method(model)
+        # Dryrun.
+        x = np.random.rand(16, *(network["input_shape"][1:]))
+        analyzer.fit(x)
+        x = np.random.rand(1, *(network["input_shape"][1:]))
+        analysis = analyzer.analyze(x)
+        self.assertEqual(tuple(analysis.shape),
+                         (1,)+tuple(network["input_shape"][1:]))
+        self.assertFalse(np.any(np.isinf(analysis.ravel())))
+        self.assertFalse(np.any(np.isnan(analysis.ravel())))
+        self.assertFalse(True)
+
+
+def test_train_analyzer(method, network_filter):
+    """Workaround for move from unit-tests to pytest."""
+    # todo: Mixing of pytest and unittest is not ideal.
+    # Move completely to pytest.
+    test_case = AnalyzerTrainTestCase(method=method,
+                                      network_filter=network_filter)
+    test_result = unittest.TextTestRunner().run(test_case)
+    assert len(test_result.errors) == 0
+    assert len(test_result.failures) == 0
+
+
+class EqualAnalyzerTestCase(BaseLayerTestCase):
+    """TestCase for analyzers execution
+
+    TestCase that applies two method to several networks and
+    runs the analyzer with random data and checks for equality
+    of the results.
+
+    :param method1: A function that returns an Analyzer class.
+    :param method2: A function that returns an Analyzer class.
+    """
+
+    def __init__(self, *args, **kwargs):
+        method1 = kwargs.pop("method1", None)
+        method2 = kwargs.pop("method2", None)
+        if method1 is not None:
+            self._method1 = method1
+        if method2 is not None:
+            self._method2 = method2
+
+        super(EqualAnalyzerTestCase, self).__init__(*args, **kwargs)
+
+    def _method1(self, model):
+        raise NotImplementedError("Set in subclass.")
+
+    def _method2(self, model):
+        raise NotImplementedError("Set in subclass.")
+
+    def _apply_test(self, network):
+        # Create model.
+        model = keras.models.Model(inputs=network["in"],
+                                   outputs=network["out"])
+        model.set_weights(_set_zero_weights_to_random(model.get_weights()))
+        # Get analyzer.
+        analyzer1 = self._method1(model)
+        analyzer2 = self._method2(model)
+        # Dryrun.
+        x = np.random.rand(1, *(network["input_shape"][1:]))*100
+        analysis1 = analyzer1.analyze(x)
+        analysis2 = analyzer2.analyze(x)
+
+        self.assertEqual(tuple(analysis1.shape),
+                         (1,)+tuple(network["input_shape"][1:]))
+        self.assertFalse(np.any(np.isinf(analysis1.ravel())))
+        self.assertFalse(np.any(np.isnan(analysis1.ravel())))
+        self.assertEqual(tuple(analysis2.shape),
+                         (1,)+tuple(network["input_shape"][1:]))
+        self.assertFalse(np.any(np.isinf(analysis2.ravel())))
+        self.assertFalse(np.any(np.isnan(analysis2.ravel())))
+
+        all_close_kwargs = {}
+        if hasattr(self, "_all_close_rtol"):
+            all_close_kwargs["rtol"] = self._all_close_rtol
+        if hasattr(self, "_all_close_atol"):
+            all_close_kwargs["atol"] = self._all_close_atol
+        #print(analysis1.sum(), analysis2.sum())
+        self.assertTrue(np.allclose(analysis1, analysis2, **all_close_kwargs))
+
+
+def test_equal_analyzer(method1, method2, network_filter):
+    """Workaround for move from unit-tests to pytest."""
+    # todo: Mixing of pytest and unittest is not ideal.
+    # Move completely to pytest.
+    test_case = EqualAnalyzerTestCase(method1=method1,
+                                      method2=method2,
+                                      network_filter=network_filter)
+    test_result = unittest.TextTestRunner().run(test_case)
+    assert len(test_result.errors) == 0
+    assert len(test_result.failures) == 0
+
+
+# todo: merge with base test case? if we don't run the analysis
+# its only half the test.
+class SerializeAnalyzerTestCase(BaseLayerTestCase):
+    """TestCase for analyzers serialization
+
+    TestCase that applies the method to several networks and
+    runs the analyzer with random data, serializes it, and
+    runs it again.
+
+    :param method: A function that returns an Analyzer class.
+    """
+
+    def __init__(self, *args, **kwargs):
+        method = kwargs.pop("method", None)
+        if method is not None:
+            self._method = method
+        super(SerializeAnalyzerTestCase, self).__init__(*args, **kwargs)
+
+    def _method(self, model):
+        raise NotImplementedError("Set in subclass.")
+
+    def _apply_test(self, network):
+        # Create model.
+        model = keras.models.Model(inputs=network["in"],
+                                   outputs=network["out"])
+        model.set_weights(_set_zero_weights_to_random(model.get_weights()))
+        # Get analyzer.
+        analyzer = self._method(model)
+        # Dryrun.
+        x = np.random.rand(1, *(network["input_shape"][1:]))
+
+        class_name, state = analyzer.save()
+        new_analyzer = AnalyzerBase.load(class_name, state)
+
+        analysis = new_analyzer.analyze(x)
+        self.assertEqual(tuple(analysis.shape),
+                         (1,)+tuple(network["input_shape"][1:]))
+        self.assertFalse(np.any(np.isinf(analysis.ravel())))
+        self.assertFalse(np.any(np.isnan(analysis.ravel())))
+
+
+def test_serialize_analyzer(method, network_filter):
+    """Workaround for move from unit-tests to pytest."""
+    # todo: Mixing of pytest and unittest is not ideal.
+    # Move completely to pytest.
+    test_case = SerializeAnalyzerTestCase(method=method,
+                                          network_filter=network_filter)
+    test_result = unittest.TextTestRunner().run(test_case)
+    assert len(test_result.errors) == 0
+    assert len(test_result.failures) == 0
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class PatternComputerTestCase(BaseLayerTestCase):
+    """TestCase pattern computation
+
+    :param method: A function that returns an PatternComputer class.
+    """
+
+    def __init__(self, *args, **kwargs):
+        method = kwargs.pop("method", None)
+        if method is not None:
+            self._method = method
+        super(PatternComputerTestCase, self).__init__(*args, **kwargs)
+
+    def _method(self, model):
+        raise NotImplementedError("Set in subclass.")
+
+    def _apply_test(self, network):
+        # Create model.
+        model = keras.models.Model(inputs=network["in"], outputs=network["out"])
+        model.set_weights(_set_zero_weights_to_random(model.get_weights()))
+        # Get computer.
+        computer = self._method(model)
+        # Dryrun.
+        x = np.random.rand(10, *(network["input_shape"][1:]))
+        computer.compute(x)
+
+
+def test_pattern_computer(method, network_filter):
+    """Workaround for move from unit-tests to pytest."""
+    # todo: Mixing of pytest and unittest is not ideal.
+    # Move completely to pytest.
+    test_case = PatternComputerTestCase(method=method,
+                                        network_filter=network_filter)
+    test_result = unittest.TextTestRunner().run(test_case)
+    assert len(test_result.errors) == 0
+    assert len(test_result.failures) == 0

+ 92 - 0
original_model/innvestigate/utils/tests/layer.py

@@ -0,0 +1,92 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import range
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.models
+import keras.engine.topology
+
+
+from ... import utils as iutils
+
+
+__all__ = [
+    "TestAnalysisHelper",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+class TestAnalysisHelper(object):
+
+    def __init__(self, model, analyzer, weights=None):
+        """ Helper class for retrieving output and analysis in test cases.
+        
+
+        :param model: A Keras layer object or a list of layer objects.
+          In this case a sequntial model will be build. The first layer
+          must have set input_shape or batch_input_shape.
+          Alternatively a tuple with input and output tensors, in which
+          case the keras modle api will be used.
+        :param analyzer: Either an analyzer class or a function
+          that takes a keras model and returns an analyzer.
+        :param weights: After creating the model set the given weights.
+        """
+
+        if isinstance(model, keras.engine.topology.Layer):
+            model = [model]
+
+        if isinstance(model, list):
+            self._model = keras.models.Sequential(model)
+        else:
+            self._model = keras.models.Model(*model)
+
+        self._input_shapes = iutils.to_list(self._model.input_shape)
+
+        if weights is not None:
+            self._model.set_weights(weights)
+
+        self._analyzer = analyzer(self._model)
+
+    @property
+    def weights(self):
+        return self._model.get_weights()
+    
+    def run(self, inputs):
+        """Runs the model given the inputs.
+
+        :return: Tuple with model output and analyzer output.
+        """
+        return_list = True
+        if not isinstance(inputs, list):
+            return_list = False
+            inputs = iutils.to_list(inputs)
+
+        augmented = []
+        for i in range(len(inputs)):
+            if len(inputs[i].shape) == len(self._input_shapes[i])-1:
+                # Augment by batch axis.
+                augmented.append(i)
+                inputs[i] = inputs[i].reshape((1,)+inputs[i].shape)
+
+        outputs = iutils.to_list(self._model.predict_on_batch(inputs))
+        analysis = iutils.to_list(self._analyzer.analyze(inputs))
+
+        for i in augmented:
+            # Remove batch axis.
+            outputs[i] = outputs[i][0]
+            analysis[i] = analysis[i][0]
+        
+        if return_list:
+            return outputs, analysis
+        else:
+            return outputs[0], analysis[0]

+ 55 - 0
original_model/innvestigate/utils/tests/networks/__init__.py

@@ -0,0 +1,55 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.backend as K
+import fnmatch
+
+from . import trivia
+from . import mnist
+from . import cifar10
+from . import imagenet
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def iterator(network_filter="*", clear_sessions=False):
+    """
+    Iterator over various networks.
+    """
+
+    def fetch_networks(module_name, module):
+        ret = [
+            ("%s.%s" % (module_name, name),
+             (module, name))
+            for name in module.__all__
+            if any((fnmatch.fnmatch(name, one_filter) or
+                    fnmatch.fnmatch("%s.%s" % (module_name, name), one_filter))
+                   for one_filter in network_filter.split(":"))
+        ]
+
+        return [x for x in sorted(ret)]
+
+    networks = (
+        fetch_networks("trivia", trivia) +
+        fetch_networks("mnist", mnist) +
+        fetch_networks("cifar10", cifar10) +
+        fetch_networks("imagenet", imagenet)
+    )
+
+    for module_name, (module, name) in networks:
+        if clear_sessions:
+            K.clear_session()
+
+        network = getattr(module, name)()
+        network["name"] = module_name
+        yield network

+ 277 - 0
original_model/innvestigate/utils/tests/networks/base.py

@@ -0,0 +1,277 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+from builtins import range
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+import keras.layers
+
+
+__all__ = [
+    "log_reg",
+
+    "mlp_2dense",
+    "mlp_3dense",
+
+    "cnn_1convb_2dense",
+    "cnn_2convb_2dense",
+    "cnn_2convb_3dense",
+    "cnn_3convb_3dense",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+# TODO: more consistent nameing
+
+def input_layer(shape, *args, **kwargs):
+    return keras.layers.Input(shape=shape[1:], *args, **kwargs)
+
+
+def dense_layer(layer_in, *args, **kwargs):
+    return keras.layers.Dense(*args, **kwargs)(layer_in)
+
+
+def conv_layer(layer_in, *args, **kwargs):
+    return keras.layers.Conv2D(*args, **kwargs)(layer_in)
+
+
+def conv_pool(layer_in, n_conv, prefix, n_filter, **kwargs):
+    conv_prefix = "%s_%%i" % prefix
+
+    ret = {}
+    current_layer = layer_in
+    for i in range(n_conv):
+        conv = conv_layer(current_layer, filters=n_filter,
+                          kernel_size=(3, 3), strides=(1, 1), padding="same",
+                          kernel_initializer="glorot_uniform", **kwargs)
+        current_layer = conv
+        ret[conv_prefix % i] = conv
+
+        ret["%s_pool" % prefix] = keras.layers.MaxPooling2D(
+            pool_size=(2, 2),
+            strides=(2, 2),
+        )(current_layer)
+    return ret
+
+
+def dropout_layer(layer_in, *args, **kwargs):
+    return keras.layers.Dropout(*args, **kwargs)(layer_in)
+
+
+def softmax(layer_in):
+    return keras.layers.Activation("softmax")(layer_in)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def log_reg(input_shape, output_n, activation=None):
+    if activation is None:
+        activation = "relu"
+
+    net = {}
+    net["in"] = input_layer(shape=input_shape)
+    net["in_flat"] = keras.layers.Flatten()(net["in"])
+    net["out"] = dense_layer(net["in_flat"], units=output_n,
+                             kernel_initializer="glorot_uniform")
+    net["sm_out"] = softmax(net["out"])
+
+    net.update({
+        "input_shape": input_shape,
+
+        "output_n": output_n,
+    })
+    return net
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def mlp_2dense(input_shape, output_n, activation=None,
+               dense_units=512, dropout_rate=0.25):
+    if activation is None:
+        activation = "relu"
+
+    net = {}
+    net["in"] = input_layer(shape=input_shape)
+    net["in_flat"] = keras.layers.Flatten()(net["in"])
+    net["dense_1"] = dense_layer(net["in_flat"], units=dense_units,
+                                 activation=activation,
+                                 kernel_initializer="glorot_uniform")
+    net["dense_1_dropout"] = dropout_layer(net["dense_1"], dropout_rate)
+    net["out"] = dense_layer(net["dense_1_dropout"], units=output_n,
+                             kernel_initializer="glorot_uniform")
+    net["sm_out"] = softmax(net["out"])
+
+    net.update({
+        "input_shape": input_shape,
+
+        "output_n": output_n,
+    })
+    return net
+
+
+def mlp_3dense(input_shape, output_n, activation=None,
+               dense_units=512, dropout_rate=0.25):
+    if activation is None:
+        activation = "relu"
+
+    net = {}
+    net["in"] = input_layer(shape=input_shape)
+    net["in_flat"] = keras.layers.Flatten()(net["in"])
+    net["dense_1"] = dense_layer(net["in_flat"], units=dense_units,
+                                 activation=activation,
+                                 kernel_initializer="glorot_uniform")
+    net["dense_1_dropout"] = dropout_layer(net["dense_1"], dropout_rate)
+    net["dense_2"] = dense_layer(net["dense_1_dropout"], units=dense_units,
+                                 activation=activation,
+                                 kernel_initializer="glorot_uniform")
+    net["dense_2_dropout"] = dropout_layer(net["dense_2"], dropout_rate)
+    net["out"] = dense_layer(net["dense_2_dropout"], units=output_n,
+                             kernel_initializer="glorot_uniform")
+    net["sm_out"] = softmax(net["out"])
+
+    net.update({
+        "input_shape": input_shape,
+
+        "output_n": output_n,
+    })
+    return net
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def cnn_1convb_2dense(input_shape, output_n, activation=None,
+                      dense_units=512, dropout_rate=0.25):
+    if activation is None:
+        activation = "relu"
+
+    net = {}
+    net["in"] = input_layer(shape=input_shape)
+    net.update(conv_pool(net["in"], 2, "conv_1", 128,
+                         activation=activation))
+    net["conv_flat"] = keras.layers.Flatten()(net["conv_1_pool"])
+    net["dense_1"] = dense_layer(net["conv_flat"], units=dense_units,
+                                 activation=activation,
+                                 kernel_initializer="glorot_uniform")
+    net["dense_1_dropout"] = dropout_layer(net["dense_1"], dropout_rate)
+    net["out"] = dense_layer(net["dense_1_dropout"], units=output_n,
+                             kernel_initializer="glorot_uniform")
+    net["sm_out"] = softmax(net["out"])
+
+    net.update({
+        "input_shape": input_shape,
+
+        "output_n": output_n,
+    })
+    return net
+
+
+def cnn_2convb_2dense(input_shape, output_n, activation=None,
+                      dense_units=512, dropout_rate=0.25):
+    if activation is None:
+        activation = "relu"
+
+    net = {}
+    net["in"] = input_layer(shape=input_shape)
+    net.update(conv_pool(net["in"], 2, "conv_1", 128,
+                         activation=activation))
+    net.update(conv_pool(net["conv_1_pool"], 2, "conv_2", 128,
+                         activation=activation))
+    net["conv_flat"] = keras.layers.Flatten()(net["conv_2_pool"])
+    net["dense_1"] = dense_layer(net["conv_flat"], units=dense_units,
+                                 activation=activation,
+                                 kernel_initializer="glorot_uniform")
+    net["dense_1_dropout"] = dropout_layer(net["dense_1"], dropout_rate)
+    net["out"] = dense_layer(net["dense_1_dropout"], units=output_n,
+                             kernel_initializer="glorot_uniform")
+    net["sm_out"] = softmax(net["out"])
+
+    net.update({
+        "input_shape": input_shape,
+
+        "output_n": output_n,
+    })
+    return net
+
+
+def cnn_2convb_3dense(input_shape, output_n, activation=None,
+                      dense_units=512, dropout_rate=0.25):
+    if activation is None:
+        activation = "relu"
+
+    net = {}
+    net["in"] = input_layer(shape=input_shape)
+    net.update(conv_pool(net["in"], 2, "conv_1", 128,
+                         activation=activation))
+    net.update(conv_pool(net["conv_1_pool"], 2, "conv_2", 128,
+                         activation=activation))
+    net["conv_flat"] = keras.layers.Flatten()(net["conv_2_pool"])
+    net["dense_1"] = dense_layer(net["conv_flat"], units=dense_units,
+                                 activation=activation,
+                                 kernel_initializer="glorot_uniform")
+    net["dense_1_dropout"] = dropout_layer(net["dense_1"], dropout_rate)
+    net["dense_2"] = dense_layer(net["dense_1_dropout"], units=dense_units,
+                                 activation=activation,
+                                 kernel_initializer="glorot_uniform")
+    net["dense_2_dropout"] = dropout_layer(net["dense_2"], dropout_rate)
+    net["out"] = dense_layer(net["dense_2_dropout"], units=output_n,
+                             kernel_initializer="glorot_uniform")
+    net["sm_out"] = softmax(net["out"])
+
+    net.update({
+        "input_shape": input_shape,
+
+        "output_n": output_n,
+    })
+    return net
+
+
+def cnn_3convb_3dense(input_shape, output_n, activation=None,
+                      dense_units=512, dropout_rate=0.25):
+    if activation is None:
+        activation = "relu"
+
+    net = {}
+    net["in"] = input_layer(shape=input_shape)
+    net.update(conv_pool(net["in"], 2, "conv_1", 128,
+                         activation=activation))
+    net.update(conv_pool(net["conv_1_pool"], 2, "conv_2", 128,
+                         activation=activation))
+    net.update(conv_pool(net["conv_2_pool"], 2, "conv_3", 128,
+                         activation=activation))
+    net["conv_flat"] = keras.layers.Flatten()(net["conv_3_pool"])
+    net["dense_1"] = dense_layer(net["conv_flat"], units=dense_units,
+                                 activation=activation,
+                                 kernel_initializer="glorot_uniform")
+    net["dense_1_dropout"] = dropout_layer(net["dense_1"], dropout_rate)
+    net["dense_2"] = dense_layer(net["dense_1_dropout"], units=dense_units,
+                                 activation=activation,
+                                 kernel_initializer="glorot_uniform")
+    net["dense_2_dropout"] = dropout_layer(net["dense_2"], dropout_rate)
+    net["out"] = dense_layer(net["dense_2_dropout"], units=output_n,
+                             kernel_initializer="glorot_uniform")
+    net["sm_out"] = softmax(net["out"])
+
+    net.update({
+        "input_shape": input_shape,
+
+        "output_n": output_n,
+    })
+    return net
+

+ 94 - 0
original_model/innvestigate/utils/tests/networks/cifar10.py

@@ -0,0 +1,94 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.backend as K
+
+from . import base
+
+
+__all__ = [
+    "log_reg",
+
+    "mlp_2dense",
+    "mlp_3dense",
+
+    "cnn_1convb_2dense",
+    "cnn_2convb_2dense",
+    "cnn_2convb_3dense",
+    "cnn_3convb_3dense",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+if K.image_data_format() == "channels_first":
+    __input_shape__ = [None, 3, 32, 32]
+else:
+    __input_shape__ = [None, 32, 32, 3]
+__output_n__ = 10
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def log_reg(activation=None):
+    return base.log_reg(__input_shape__, __output_n__,
+                        activation=activation)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def mlp_2dense(activation=None):
+    return base.mlp_2dense(__input_shape__, __output_n__,
+                           activation=activation,
+                           dense_units=1024, dropout_rate=0.5)
+
+
+def mlp_3dense(activation=None):
+    return base.mlp_3dense(__input_shape__, __output_n__,
+                           activation=activation,
+                           dense_units=1024, dropout_rate=0.5)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def cnn_1convb_2dense(activation=None):
+    return base.cnn_1convb_2dense(__input_shape__, __output_n__,
+                                  activation=activation,
+                                  dense_units=1024, dropout_rate=0.5)
+
+
+def cnn_2convb_2dense(activation=None):
+    return base.cnn_2convb_2dense(__input_shape__, __output_n__,
+                                  activation=activation,
+                                  dense_units=1024, dropout_rate=0.5)
+
+
+def cnn_2convb_3dense(activation=None):
+    return base.cnn_2convb_3dense(__input_shape__, __output_n__,
+                                  activation=activation,
+                                  dense_units=1024, dropout_rate=0.5)
+
+
+def cnn_3convb_3dense(activation=None):
+    return base.cnn_3convb_3dense(__input_shape__, __output_n__,
+                                  activation=activation,
+                                  dense_units=1024, dropout_rate=0.5)

+ 210 - 0
original_model/innvestigate/utils/tests/networks/imagenet.py

@@ -0,0 +1,210 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.backend as K
+import keras.layers
+import numpy as np
+import warnings
+
+from . import base
+from . import mnist
+from ....applications import imagenet
+
+__all__ = [
+    "vgg16_custom",
+    "vgg16",
+    "vgg19",
+    "resnet50",
+    "inception_v3",
+    "inception_resnet_v2",
+    "densenet121",
+    "densenet169",
+    "densenet201",
+    "nasnet_large",
+    "nasnet_mobile",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+VGG16_OFFSET = np.array([103.939, 116.779, 123.68])
+
+
+def vgg16_custom_preprocess(X):
+    import innvestigate.utils.visualizations as ivis
+    X = ivis.preprocess_images(X, color_coding="RGBtoBGR")
+
+    if X.shape[1] == 3:
+        shape = [1, 3, 1, 1]
+    else:
+        shape = [1, 1, 1, 3]
+
+    offset = VGG16_OFFSET.reshape(shape)
+    # Remove pixel-wise mean.
+    X -= offset
+    return X
+
+
+def vgg16_custom(activation=None):
+    if activation is None:
+        activation = "relu"
+
+    if K.image_data_format() == "channels_first":
+        input_shape = [None, 3, 224, 224]
+    else:
+        input_shape = [None, 224, 224, 3]
+    output_n = 1000
+
+    net = {}
+    net["in"] = base.input_layer(shape=input_shape)
+
+    net.update(base.conv_pool(
+        net["in"], 2, "conv_1", 64,
+        activation=activation,
+    ))
+    net.update(base.conv_pool(
+        net["conv_1_pool"], 2, "conv_2", 128,
+        activation=activation,
+    ))
+    net.update(base.conv_pool(
+        net["conv_2_pool"], 3, "conv_3", 256,
+        activation=activation,
+    ))
+    net.update(base.conv_pool(
+        net["conv_3_pool"], 3, "conv_4", 512,
+        activation=activation,
+    ))
+    net.update(base.conv_pool(
+        net["conv_4_pool"], 3, "conv_5", 512,
+        activation=activation,
+    ))
+
+    net["conv_flat"] = keras.layers.Flatten()(net["conv_5_pool"])
+    net["dense_1"] = base.dense_layer(net["conv_flat"], units=4096,
+                                      activation=activation,
+                                      kernel_initializer="glorot_uniform")
+    net["dense_1_dropout"] = base.dropout_layer(net["dense_1"], 0.5)
+    net["dense_2"] = base.dense_layer(net["dense_1_dropout"], units=4096,
+                                      activation=activation,
+                                      kernel_initializer="glorot_uniform")
+    net["dense_2_dropout"] = base.dropout_layer(net["dense_2"], 0.5)
+    net["out"] = base.dense_layer(net["dense_2_dropout"], units=output_n,
+                                  kernel_initializer="glorot_uniform")
+    net["sm_out"] = base.softmax(net["out"])
+
+    net.update({
+        "input_shape": input_shape,
+        "preprocess_f": vgg16_custom_preprocess,
+        "output_n": output_n,
+    })
+    return net
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def vgg16():
+    ret = imagenet.vgg16()
+    ret["output_n"] = 1000
+    return ret
+
+
+def vgg19():
+    ret = imagenet.vgg19()
+    ret["output_n"] = 1000
+    return ret
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def resnet50():
+    ret = imagenet.resnet50()
+    ret["output_n"] = 1000
+    return ret
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def inception_v3():
+    ret = imagenet.inception_v3()
+    ret["output_n"] = 1000
+    return ret
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def inception_resnet_v2():
+    ret = imagenet.inception_resnet_v2()
+    ret["output_n"] = 1000
+    return ret
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def densenet121():
+    ret = imagenet.densenet121()
+    ret["output_n"] = 1000
+    return ret
+
+
+def densenet169():
+    ret = imagenet.densenet169()
+    ret["output_n"] = 1000
+    return ret
+
+
+def densenet201():
+    ret = imagenet.densenet201()
+    ret["output_n"] = 1000
+    return ret
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def nasnet_large():
+    if K.image_data_format() == "channels_first":
+        warnings.warn("NASNet is not available for channels first. "
+                      "Return dummy net.")
+        return mnist.log_reg()
+
+    ret = imagenet.nasnet_large()
+    ret["output_n"] = 1000
+    return ret
+
+
+def nasnet_mobile():
+    if K.image_data_format() == "channels_first":
+        warnings.warn("NASNet is not available for channels first. "
+                      "Return dummy net.")
+        return mnist.log_reg()
+
+    ret = imagenet.nasnet_mobile()
+    ret["output_n"] = 1000
+    return ret

+ 94 - 0
original_model/innvestigate/utils/tests/networks/mnist.py

@@ -0,0 +1,94 @@
+# Get Python six functionality:
+from __future__ import\
+    absolute_import, print_function, division, unicode_literals
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+import keras.backend as K
+
+from . import base
+
+
+__all__ = [
+    "log_reg",
+
+    "mlp_2dense",
+    "mlp_3dense",
+
+    "cnn_1convb_2dense",
+    "cnn_2convb_2dense",
+    "cnn_2convb_3dense",
+    "cnn_3convb_3dense",
+]
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+if K.image_data_format() == "channels_first":
+    __input_shape__ = [None, 1, 28, 28]
+else:
+    __input_shape__ = [None, 28, 28, 1]
+__output_n__ = 10
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def log_reg(activation=None):
+    return base.log_reg(__input_shape__, __output_n__,
+                        activation=activation)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def mlp_2dense(activation=None):
+    return base.mlp_2dense(__input_shape__, __output_n__,
+                           activation=activation,
+                           dense_units=512, dropout_rate=0.25)
+
+
+def mlp_3dense(activation=None):
+    return base.mlp_3dense(__input_shape__, __output_n__,
+                           activation=activation,
+                           dense_units=512, dropout_rate=0.25)
+
+
+###############################################################################
+###############################################################################
+###############################################################################
+
+
+def cnn_1convb_2dense(activation=None):
+    return base.cnn_1convb_2dense(__input_shape__, __output_n__,
+                                  activation=activation,
+                                  dense_units=512, dropout_rate=0.25)
+
+
+def cnn_2convb_2dense(activation=None):
+    return base.cnn_2convb_2dense(__input_shape__, __output_n__,
+                                  activation=activation,
+                                  dense_units=512, dropout_rate=0.25)
+
+
+def cnn_2convb_3dense(activation=None):
+    return base.cnn_2convb_3dense(__input_shape__, __output_n__,
+                                  activation=activation,
+                                  dense_units=512, dropout_rate=0.25)
+
+
+def cnn_3convb_3dense(activation=None):
+    return base.cnn_3convb_3dense(__input_shape__, __output_n__,
+                                  activation=activation,
+                                  dense_units=512, dropout_rate=0.25)

Some files were not shown because too many files changed in this diff