{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":20270,"databundleVersionId":1222630,"sourceType":"competition"},{"sourceId":1295659,"sourceType":"datasetVersion","datasetId":749200},{"sourceId":1322494,"sourceType":"datasetVersion","datasetId":688574}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<center><img src='https://raw.githubusercontent.com/dimitreOliveira/MachineLearning/master/Kaggle/SIIM-ISIC%20Melanoma%20Classification/banner.png' height=\"350\"></center>\n<p>\n<h1><center> SIIM-ISIC Melanoma Classification </center></h1>\n<h2><center> Melanoma Classification - SHAP model explained </center></h2>\n<p>\n\n#### About SHAP from [the reposiroty](https://github.com/slundberg/shap)\n<center><img src='https://raw.githubusercontent.com/slundberg/shap/master/docs/artwork/shap_header.png' width=\"500\" height=\"150\"></center>\n\n#### SHAP (SHapley Additive exPlanations) is a game theoretic approach to explain the output of any machine learning model. It connects optimal credit allocation with local explanations using the classic Shapley values from game theory and their related extensions.","metadata":{}},{"cell_type":"markdown","source":"## Dependencies","metadata":{}},{"cell_type":"code","source":"!pip install --quiet image-classifiers\n\nimport warnings, json, re, glob, math, shutil, os, shap\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport tensorflow as tf\nimport tensorflow.keras.layers as L\nimport tensorflow.keras.backend as K\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras import Model\nfrom kaggle_datasets import KaggleDatasets\nfrom sklearn.model_selection import KFold\nfrom classification_models.tfkeras import Classifiers\n\nSEED = 0\nwarnings.filterwarnings(\"ignore\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-10-24T18:35:53.830748Z","iopub.execute_input":"2024-10-24T18:35:53.831144Z","iopub.status.idle":"2024-10-24T18:36:26.074240Z","shell.execute_reply.started":"2024-10-24T18:35:53.831104Z","shell.execute_reply":"2024-10-24T18:36:26.073207Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load data","metadata":{}},{"cell_type":"code","source":"database_base_path = '/kaggle/input/siim-isic-melanoma-classification/'\ntrain = pd.read_csv(database_base_path + 'train.csv')\ntest = pd.read_csv(database_base_path + 'test.csv')\n\nprint('Train samples: %d' % len(train))\ndisplay(train.head())\nprint(f'Test samples: {len(test)}')\ndisplay(test.head())\n\n# pre-process data\ntrain['image_name'] = train['image_name'].apply(lambda x: x + '.jpg')\ntrain['target'] = train['target'].astype(str)\n\nGCS_PATH = KaggleDatasets().get_gcs_path('melanoma-256x256')\nTRAINING_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/train*.tfrec')\nTEST_FILENAMES = tf.io.gfile.glob(GCS_PATH + '/test*.tfrec')","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-10-24T18:36:26.076202Z","iopub.execute_input":"2024-10-24T18:36:26.076846Z","iopub.status.idle":"2024-10-24T18:36:27.705505Z","shell.execute_reply.started":"2024-10-24T18:36:26.076805Z","shell.execute_reply":"2024-10-24T18:36:27.704481Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Parameters","metadata":{}},{"cell_type":"code","source":"HEIGHT = 256\nWIDTH = 256\nCHANNELS = 3\nBATCH_SIZE = 64\nAUTO = tf.data.experimental.AUTOTUNE\n\n# SHAP parameters\nimages_to_explain = ['ISIC_0074311.jpg', 'ISIC_0074542.jpg', 'ISIC_0075663.jpg', 'ISIC_0075914.jpg', \n                     'ISIC_0076262.jpg', 'ISIC_0082543.jpg', 'ISIC_0082934.jpg', 'ISIC_0083035.jpg', \n                     'ISIC_0084086.jpg', 'ISIC_0084270.jpg', 'ISIC_0149568.jpg', 'ISIC_0188432.jpg', \n                     'ISIC_0207268.jpg', 'ISIC_0232101.jpg', 'ISIC_0247330.jpg', 'ISIC_0528044.jpg', \n                     'ISIC_1219894.jpg', 'ISIC_2776906.jpg']\n\neval_df = train[train['image_name'].isin(images_to_explain)]\n\nos.makedirs('to_explain/')\nfor filename in images_to_explain:\n    shutil.copy('/kaggle/input/siim-isic-melanoma-classification/jpeg/train/' + filename, 'to_explain/')","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2024-10-24T18:36:27.707784Z","iopub.execute_input":"2024-10-24T18:36:27.708122Z","iopub.status.idle":"2024-10-24T18:36:28.027972Z","shell.execute_reply.started":"2024-10-24T18:36:27.708086Z","shell.execute_reply":"2024-10-24T18:36:28.026961Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Auxiliary functions","metadata":{}},{"cell_type":"code","source":"# Datasets utility functions\nUNLABELED_TFREC_FORMAT = {\n    \"image\": tf.io.FixedLenFeature([], tf.string), # tf.string means bytestring\n    \"image_name\": tf.io.FixedLenFeature([], tf.string), # shape [] means single element\n    # meta features\n    \"patient_id\": tf.io.FixedLenFeature([], tf.int64),\n    \"sex\": tf.io.FixedLenFeature([], tf.int64),\n    \"age_approx\": tf.io.FixedLenFeature([], tf.int64),\n    \"anatom_site_general_challenge\": tf.io.FixedLenFeature([], tf.int64),\n}\n\ndef decode_image(image_data, height, width, channels):\n    image = tf.image.decode_jpeg(image_data, channels=channels)\n    image = tf.cast(image, tf.float32) / 255.0\n    image = tf.reshape(image, [height, width, channels])\n    return image\n\n# Test function\ndef read_unlabeled_tfrecord(example, height=HEIGHT, width=WIDTH, channels=CHANNELS):\n    example = tf.io.parse_single_example(example, UNLABELED_TFREC_FORMAT)\n    image = decode_image(example['image'], height, width, channels)\n    image_name = example['image_name']\n    # meta features\n    data = {}\n    data['patient_id'] = tf.cast(example['patient_id'], tf.int32)\n    data['sex'] = tf.cast(example['sex'], tf.int32)\n    data['age_approx'] = tf.cast(example['age_approx'], tf.int32)\n    data['anatom_site_general_challenge'] = tf.cast(tf.one_hot(example['anatom_site_general_challenge'], 7), tf.int32)\n    \n    return {'input_image': image, 'input_tabular': data}, image_name # returns a dataset of (image, data, image_name)\n\ndef load_dataset_test(filenames, buffer_size=-1):\n    dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=buffer_size) # automatically interleaves reads from multiple files\n    dataset = dataset.map(read_unlabeled_tfrecord, num_parallel_calls=buffer_size)\n    # returns a dataset of (image, data, label, image_name) pairs if labeled=True or (image, data, image_name) pairs if labeled=False\n    return dataset\n\ndef get_test_dataset(filenames, batch_size=32, buffer_size=-1):\n    dataset = load_dataset_test(filenames, buffer_size=buffer_size)\n    dataset = dataset.batch(batch_size, drop_remainder=False)\n    dataset = dataset.prefetch(buffer_size)\n    return dataset\n\n# Custom SHAP plot\ndef image_plot(shap_values, pixel_values, labels=None, preds=None, names=None, width=20, aspect=0.2, hspace=0.2, labelpad=None, show=True, fig_size=None):\n    \"\"\" Plots SHAP values for image inputs.\n    Parameters\n    ----------\n    shap_values : [numpy.array]\n        List of arrays of SHAP values. Each array has the shap (# samples x width x height x channels), and the\n        length of the list is equal to the number of model outputs that are being explained.\n    pixel_values : numpy.array\n        Matrix of pixel values (# samples x width x height x channels) for each image. It should be the same\n        shape as each array in the shap_values list of arrays.\n    labels : list\n        List of names for each of the model outputs that are being explained. This list should be the same length\n        as the shap_values list.\n    width : float\n        The width of the produced matplotlib plot.\n    labelpad : float\n        How much padding to use around the model output labels.\n    show : bool\n        Whether matplotlib.pyplot.show() is called before returning. Setting this to False allows the plot\n        to be customized further after it has been created.\n    \"\"\"\n\n    multi_output = True\n    if type(shap_values) != list:\n        multi_output = False\n        shap_values = [shap_values]\n\n    # make sure labels\n    if labels is not None:\n        assert labels.shape[0] == shap_values[0].shape[0], \"Labels must have same row count as shap_values arrays!\"\n        if multi_output:\n            assert labels.shape[1] == len(shap_values), \"Labels must have a column for each output in shap_values!\"\n        else:\n            assert len(labels.shape) == 1, \"Labels must be a vector for single output shap_values.\"\n\n    label_kwargs = {} if labelpad is None else {'pad': labelpad}\n\n    # plot our explanations\n    x = pixel_values\n    if fig_size is None:\n        fig_size = np.array([3 * (len(shap_values) + 1), 2.5 * (x.shape[0] + 1)])\n        if fig_size[0] > width:\n            fig_size *= width / fig_size[0]\n    fig, axes = plt.subplots(nrows=x.shape[0], ncols=len(shap_values) + 1, figsize=fig_size)\n    if len(axes.shape) == 1:\n        axes = axes.reshape(1,axes.size)\n    for row in range(x.shape[0]):\n        x_curr = x[row].copy()\n\n        # make sure\n        if len(x_curr.shape) == 3 and x_curr.shape[2] == 1:\n            x_curr = x_curr.reshape(x_curr.shape[:2])\n        if x_curr.max() > 1:\n            x_curr /= 255.\n\n        # get a grayscale version of the image\n        if len(x_curr.shape) == 3 and x_curr.shape[2] == 3:\n            x_curr_gray = (0.2989 * x_curr[:,:,0] + 0.5870 * x_curr[:,:,1] + 0.1140 * x_curr[:,:,2]) # rgb to gray\n        else:\n            x_curr_gray = x_curr\n\n        axes[row,0].imshow(x_curr, cmap=plt.get_cmap('gray'))\n        axes[row,0].set_title(f'Image: {names[row]}', **label_kwargs)\n        axes[row,0].axis('off')\n        if len(shap_values[0][row].shape) == 2:\n            abs_vals = np.stack([np.abs(shap_values[i]) for i in range(len(shap_values))], 0).flatten()\n        else:\n            abs_vals = np.stack([np.abs(shap_values[i].sum(-1)) for i in range(len(shap_values))], 0).flatten()\n        max_val = np.nanpercentile(abs_vals, 99.9)\n        for i in range(len(shap_values)):\n            if labels is not None:\n                axes[row,i+1].set_title(f'Label: {labels[row,i]} Pred: {preds[row,i]:.2f}', **label_kwargs)\n            sv = shap_values[i][row] if len(shap_values[i][row].shape) == 2 else shap_values[i][row].sum(-1)\n            axes[row,i+1].imshow(x_curr_gray, cmap=plt.get_cmap('gray'), alpha=0.15, extent=(-1, sv.shape[1], sv.shape[0], -1))\n            im = axes[row,i+1].imshow(sv, cmap=shap.plots.colors.red_transparent_blue, vmin=-max_val, vmax=max_val)\n            axes[row,i+1].axis('off')\n    if hspace == 'auto':\n        fig.tight_layout()\n    else:\n        fig.subplots_adjust(hspace=hspace)\n    cb = fig.colorbar(im, ax=np.ravel(axes).tolist(), label=\"SHAP value\", orientation=\"horizontal\", aspect=fig_size[0]/aspect)\n    cb.outline.set_visible(False)\n    if show:\n        plt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-10-24T18:36:28.030244Z","iopub.execute_input":"2024-10-24T18:36:28.030589Z","iopub.status.idle":"2024-10-24T18:36:28.062692Z","shell.execute_reply.started":"2024-10-24T18:36:28.030556Z","shell.execute_reply":"2024-10-24T18:36:28.061605Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model\n\n### I will be using a ResNet18 model, trained with 3-Fold data.","metadata":{}},{"cell_type":"code","source":"def model_fn(input_shape):\n    input_image = L.Input(shape=input_shape, name='input_image')\n    ResNet18, preprocess_input = Classifiers.get('resnet18')\n    base_model = ResNet18(input_shape=input_shape, \n                          weights=None, \n                          include_top=False)\n\n    x = base_model(input_image)\n    x = L.GlobalAveragePooling2D()(x)\n    output = L.Dense(1, activation='sigmoid')(x)\n    \n    model = Model(inputs=input_image, outputs=output)\n    \n    return model\n\nmodel = model_fn((HEIGHT, WIDTH, CHANNELS))\nmodel.load_weights('/kaggle/input/shap-model/ResNet_18.h5') # load pre-trained weights\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2024-10-24T18:36:28.063951Z","iopub.execute_input":"2024-10-24T18:36:28.064292Z","iopub.status.idle":"2024-10-24T18:36:29.703372Z","shell.execute_reply.started":"2024-10-24T18:36:28.064259Z","shell.execute_reply":"2024-10-24T18:36:29.702434Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# First lets make the evaluated set predictions","metadata":{}},{"cell_type":"code","source":"# data generator\neval_datagen = ImageDataGenerator(rescale=1./255)\n\neval_generator=eval_datagen.flow_from_dataframe(\n    dataframe=eval_df,\n    directory='to_explain/',\n    x_col='image_name',\n    y_col='target',\n    class_mode='binary', \n    batch_size=BATCH_SIZE,   \n    target_size=(HEIGHT, WIDTH),\n    shuffle=False,\n    seed=SEED)\n\n# add predictions\neval_df['preds'] = model.predict(eval_generator)\ndisplay(eval_df.head())","metadata":{"execution":{"iopub.status.busy":"2024-10-24T18:36:29.704611Z","iopub.execute_input":"2024-10-24T18:36:29.704953Z","iopub.status.idle":"2024-10-24T18:36:36.705696Z","shell.execute_reply.started":"2024-10-24T18:36:29.704919Z","shell.execute_reply":"2024-10-24T18:36:36.704750Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# SHAP model explainability\n\n### For this experiment we will use `GradientExplainer`, below is an example from the Imagenet dataset applied to a VGG16 model.\n\n#### The more `pink` each pixel is the more it contributes to the image being classified on a specific class, and the more `blue` it is more the pixel contributes for it not being of that class. \n\n![](https://raw.githubusercontent.com/slundberg/shap/master/docs/artwork/gradient_imagenet_plot.png)\n\n#### Above for the 1st image the beak and wings of the dowitcher had a big contribution on assigning it to the correct class, and for the 2nd image the face of the meerkat contributed a lot to correctly classify it, for those examples we can assume that the model is doing a good job.","metadata":{}},{"cell_type":"markdown","source":"# Below are the images that will be explained by SHAP, you can also see the label and the model's prediction for each image.","metadata":{}},{"cell_type":"code","source":"n_explain = 18\neval_generator.batch_size = n_explain # background dataset\nbackground, lbls = next(eval_generator)\nlbls = lbls.reshape(lbls.shape[0], 1)\n\nfig, axes = plt.subplots(6, 3, figsize=(20, 14))\naxes = axes.flatten()\nfor x in range(6):\n    axes[x].imshow(background[x])\n    axes[x+6].imshow(background[x+6])\n    axes[x+12].imshow(background[x+12])\n    \n    axes[x].set_title(f\"Image {eval_df['image_name'].values[x]}, Label: {eval_df['target'].values[x]}, Pred: {eval_df['preds'].values[x]:.2f}\")\n    axes[x+6].set_title(f\"Image {eval_df['image_name'].values[x+6]}, Label: {eval_df['target'].values[x+6]}, Pred: {eval_df['preds'].values[x+6]:.2f}\")\n    axes[x+12].set_title(f\"Image {eval_df['image_name'].values[x+12]}, Label: {eval_df['target'].values[x+12]}, Pred: {eval_df['preds'].values[x+12]:.2f}\")\n    \n    axes[x].set_axis_off()\n    axes[x+6].set_axis_off()\n    axes[x+12].set_axis_off()\n    \nplt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-10-24T18:36:36.706965Z","iopub.execute_input":"2024-10-24T18:36:36.707262Z","iopub.status.idle":"2024-10-24T18:36:40.071712Z","shell.execute_reply.started":"2024-10-24T18:36:36.707231Z","shell.execute_reply":"2024-10-24T18:36:40.070492Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"As we can see the model missed all the predictions that the label was 1, we will also see the explanation for those images.","metadata":{}},{"cell_type":"markdown","source":"# Using SHAP to explain the images","metadata":{}},{"cell_type":"code","source":"# explain predictions of the model on \"n_explain\" images\ne = shap.GradientExplainer(model, background)\nshap_values = e.shap_values(background)\n\n# plot the feature attributions\nimage_plot(shap_values, background[:10], labels=lbls, preds=eval_df['preds'].values[:10].reshape(10, 1), names=eval_df['image_name'].values[:10], \n           hspace=0.2, fig_size=(20, 40))","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-10-24T18:36:40.073116Z","iopub.execute_input":"2024-10-24T18:36:40.073812Z","iopub.status.idle":"2024-10-24T18:36:42.546262Z","shell.execute_reply.started":"2024-10-24T18:36:40.073767Z","shell.execute_reply":"2024-10-24T18:36:42.544678Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"At all the above images the prediction score was very low, which resulted in the explained images (right images) being mostly gray, if you look really close, you may see some pink dots. At the same time, we can notice that at least the model pays attention mainly to the skin marks.\n\nOne interesting thing here is that the model seems to do not care about hair or the mm scale.","metadata":{}},{"cell_type":"markdown","source":"# Now the explanation for the positive images","metadata":{}},{"cell_type":"code","source":"shap_values_positive = shap.GradientExplainer(model, background[10:15]).shap_values(background[10:15])\n\n# plot the feature attributions\nimage_plot(shap_values_positive, background[10:15], labels=lbls[10:15], preds=eval_df['preds'].values[10:15].reshape(5, 1), \n           names=eval_df['image_name'].values[10:15], fig_size=(20, 20))","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-10-24T18:36:46.873573Z","iopub.execute_input":"2024-10-24T18:36:46.874387Z","iopub.status.idle":"2024-10-24T18:36:47.764485Z","shell.execute_reply.started":"2024-10-24T18:36:46.874345Z","shell.execute_reply":"2024-10-24T18:36:47.763061Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Here on the positive images you can see much more pink dots, they mean that those pixels contributed to the class 1 prediction, even if the model did not predicted the images as class 1, some part of the images weighted on that direction.\n\n- In this 1st image the model really messed, it gave attention to a part of the image that did not even had a skin mark.\n- At the 4th image (ISIC_0232101) the model gave some attention to the mm scale, and also scattered some attention across the image, not focussing on the skin marg too much.\n\n### Looking at those images it is clear that the model has a lot to learn about the positive classes.\n\n#### If you liked this experiment leave a comment below I may update it with a better model or make another version with both image and tabular data.","metadata":{}},{"cell_type":"markdown","source":"# Images that were predicted as postive or were very close","metadata":{}},{"cell_type":"code","source":"shap_values_positive = shap.GradientExplainer(model, background[-3:]).shap_values(background[-3:])\n\n# plot the feature attributions\nimage_plot(shap_values_positive, background[-3:], labels=lbls[-3:], preds=eval_df['preds'].values[-3:].reshape(3, 1), \n           names=eval_df['image_name'].values[-3:], fig_size=(20, 20))","metadata":{"execution":{"iopub.status.busy":"2024-10-24T18:36:42.549016Z","iopub.status.idle":"2024-10-24T18:36:42.549383Z","shell.execute_reply.started":"2024-10-24T18:36:42.549196Z","shell.execute_reply":"2024-10-24T18:36:42.549214Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Find the XAI Scores","metadata":{}},{"cell_type":"code","source":"import os\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport shap\nimport logging\nfrom typing import List, Dict, Any, Optional\nfrom sklearn.metrics import mean_squared_error\n\n# Set up logging\nlogging.basicConfig(level=logging.INFO)\nlogger = logging.getLogger(__name__)\n\ndef get_model_input_shape(model) -> tuple:\n    \"\"\"\n    Get the expected input shape from the model\n    \"\"\"\n    try:\n        input_shape = model.layers[0].input_shape\n        if isinstance(input_shape, list):\n            input_shape = input_shape[0]\n        return (input_shape[1], input_shape[2])\n    except Exception as e:\n        logger.error(f\"Error getting model input shape: {str(e)}\")\n        raise\n\ndef get_image_path(image_name: str) -> str:\n    \"\"\"\n    Get the full path for an image\n    \"\"\"\n    return os.path.join(\n        '/kaggle/input/siim-isic-melanoma-classification/jpeg/test',\n        f'{image_name}.jpg'\n    )\n\ndef load_and_preprocess_image(image_path: str, target_size: tuple) -> np.ndarray:\n    \"\"\"\n    Load and preprocess an image\n    \"\"\"\n    if not os.path.exists(image_path):\n        raise FileNotFoundError(f\"Image file not found: {image_path}\")\n        \n    img = cv2.imread(image_path, cv2.IMREAD_COLOR)\n    if img is None:\n        raise ValueError(f\"Failed to load image: {image_path}\")\n        \n    # Convert BGR to RGB\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    \n    # Resize to target size\n    img = cv2.resize(img, target_size)\n    img = img.astype(np.float32) / 255.0\n    \n    return img\n\ndef evaluate_shap_explanation(\n    model,\n    image: np.ndarray,\n    shap_values,\n    background_value: float = 0.0\n) -> Dict[str, float]:\n    \"\"\"\n    Evaluate SHAP explanation with metrics similar to LIME evaluation\n    \"\"\"\n    try:\n        # Ensure proper shapes\n        if len(image.shape) == 3:\n            image = np.expand_dims(image, axis=0)\n        \n        original_pred = model.predict(image)\n        if len(original_pred.shape) > 1:\n            original_pred = original_pred[0]\n            \n        # Calculate fidelity\n        shap_magnitude = np.abs(shap_values[0]).mean(axis=-1)\n        flat_magnitude = shap_magnitude.flatten()\n        top_pixels = np.argsort(flat_magnitude)[-100:]\n        \n        perturbed_image = image.copy()\n        for idx in top_pixels:\n            row, col = np.unravel_index(idx, shap_magnitude.shape)\n            perturbed_image[0, row, col] = background_value\n            \n        perturbed_pred = model.predict(perturbed_image)\n        if len(perturbed_pred.shape) > 1:\n            perturbed_pred = perturbed_pred[0]\n            \n        fidelity = np.mean(np.abs(original_pred - perturbed_pred))\n        \n        # Calculate unambiguity\n        unambiguity = np.var(shap_values[0])\n        \n        # Calculate interpretability\n        threshold = 0.1 * np.abs(shap_values[0]).max()\n        significant_features = np.sum(np.abs(shap_values[0]) > threshold)\n        total_features = shap_values[0].size\n        interpretability = 1.0 - (significant_features / total_features)\n        \n        return {\n            'fidelity': float(fidelity),\n            'unambiguity': float(unambiguity),\n            'interpretability': float(interpretability)\n        }\n        \n    except Exception as e:\n        logger.error(f\"Error in evaluate_shap_explanation: {str(e)}\")\n        raise\n\ndef evaluate_shap_explanations(\n    image_names: List[str],\n    model,\n    background_dataset: Optional[np.ndarray] = None,\n    num_background: int = 100\n) -> pd.DataFrame:\n    \"\"\"\n    Evaluate SHAP explanations for multiple images\n    \"\"\"\n    results = []\n    \n    # Get the expected input shape from the model\n    try:\n        input_shape = get_model_input_shape(model)\n        logger.info(f\"Model expects input shape: {input_shape}\")\n    except Exception as e:\n        logger.error(f\"Could not determine model input shape: {str(e)}\")\n        raise\n    \n    # Initialize SHAP explainer\n    try:\n        if background_dataset is None:\n            # Create simple background distribution\n            background = np.zeros((num_background,) + input_shape + (3,))\n        else:\n            # Ensure background dataset has correct shape\n            if background_dataset.shape[1:3] != input_shape:\n                resized_background = []\n                for img in background_dataset:\n                    resized = cv2.resize(img, input_shape)\n                    resized_background.append(resized)\n                background_dataset = np.array(resized_background)\n            \n            if len(background_dataset) > num_background:\n                indices = np.random.choice(len(background_dataset), num_background, replace=False)\n                background = background_dataset[indices]\n            else:\n                background = background_dataset\n                \n        explainer = shap.DeepExplainer(model, background)\n        logger.info(\"SHAP explainer initialized successfully\")\n        \n    except Exception as e:\n        logger.error(f\"Error initializing SHAP explainer: {str(e)}\")\n        raise\n    \n    for image_name in image_names:\n        logger.info(f\"Processing image: {image_name}\")\n        try:\n            # Get image path\n            image_path = get_image_path(image_name)\n            \n            # Load and preprocess image\n            img = load_and_preprocess_image(image_path, input_shape)\n            \n            # Expand dimensions for batch\n            img_batch = np.expand_dims(img, axis=0)\n            \n            logger.info(f\"Image shape after preprocessing: {img_batch.shape}\")\n            \n            # Calculate SHAP values\n            try:\n                shap_values = explainer.shap_values(img_batch)\n                logger.info(f\"SHAP values calculated successfully for {image_name}\")\n                \n                # Calculate metrics\n                metrics = evaluate_shap_explanation(model, img_batch, shap_values)\n                metrics['image_name'] = image_name\n                results.append(metrics)\n                \n            except Exception as e:\n                logger.error(f\"Error in SHAP calculation: {str(e)}\")\n                raise\n                \n        except Exception as e:\n            logger.error(f\"Error processing image {image_name}: {str(e)}\")\n            results.append({\n                'image_name': image_name,\n                'fidelity': None,\n                'unambiguity': None,\n                'interpretability': None,\n                'error': str(e)\n            })\n    \n    return pd.DataFrame(results)\n\ndef compare_explanations(\n    image_names: List[str],\n    model,\n    segment_fn,\n    background_dataset: Optional[np.ndarray] = None\n) -> pd.DataFrame:\n    \"\"\"\n    Compare LIME and SHAP explanations for the same set of images\n    \"\"\"\n    # Get model input shape\n    input_shape = get_model_input_shape(model)\n    \n    # Get LIME evaluations\n    lime_results = evaluate_explanations(image_names, model, segment_fn, image_size=input_shape)\n    \n    # Get SHAP evaluations\n    shap_results = evaluate_shap_explanations(image_names, model, background_dataset)\n    \n    # Combine results\n    combined_results = pd.merge(\n        lime_results.add_prefix('lime_'),\n        shap_results.add_prefix('shap_'),\n        left_on='lime_image_name',\n        right_on='shap_image_name'\n    )\n    \n    return combined_results\n\ndef verify_dataset_path() -> bool:\n    \"\"\"\n    Verify the dataset structure and availability\n    \"\"\"\n    dataset_path = '/kaggle/input/siim-isic-melanoma-classification/jpeg/test'\n    \n    if not os.path.exists(dataset_path):\n        logger.error(f\"Dataset directory not found: {dataset_path}\")\n        return False\n        \n    try:\n        files = os.listdir(dataset_path)\n        logger.info(f\"Found {len(files)} files in dataset directory\")\n        if len(files) > 0:\n            logger.info(f\"Example files: {files[:3]}\")\n        return True\n    except Exception as e:\n        logger.error(f\"Error accessing dataset directory: {str(e)}\")\n        return False","metadata":{"execution":{"iopub.status.busy":"2024-10-24T18:45:30.896400Z","iopub.execute_input":"2024-10-24T18:45:30.896882Z","iopub.status.idle":"2024-10-24T18:45:30.931074Z","shell.execute_reply.started":"2024-10-24T18:45:30.896840Z","shell.execute_reply":"2024-10-24T18:45:30.930077Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# For just SHAP evaluation\nimage_names = ['ISIC_0052060']\nresults_df = evaluate_shap_explanations(image_names, model)\nprint(\"\\nSHAP Evaluation Metrics:\")\nprint(results_df)","metadata":{"execution":{"iopub.status.busy":"2024-10-24T18:45:37.771278Z","iopub.execute_input":"2024-10-24T18:45:37.771977Z","iopub.status.idle":"2024-10-24T18:45:37.909131Z","shell.execute_reply.started":"2024-10-24T18:45:37.771936Z","shell.execute_reply":"2024-10-24T18:45:37.907909Z"},"trusted":true},"outputs":[],"execution_count":null}]}