{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":6799,"databundleVersionId":4225553},{"sourceType":"datasetVersion","sourceId":11439370,"datasetId":7095751,"databundleVersionId":11878886}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# How to use LIME-XAI method to explain Vision models - eXplainable Computer Vision (XCV)\n- Get basic familier with [***LIME***](https://github.com/marcotcr/lime) at [**this video**](https://www.youtube.com/watch?v=CY3t11vuuOM&t=1343s)\n- Learn about How LIME-Stratified work here at [**this article**](https://muhammad-rashid.medium.com/stratified-lime-to-generate-image-explanation-an-improved-version-of-lime-image-6b9668f03f1f)\n- LIME Image codes on [***GitHub***](https://github.com/rashidrao-pk/lime_stratified)\n- LIME Image Examples on [***GitHub***](https://github.com/rashidrao-pk/lime-stratified-examples)\n- Dataset used: [***ImageNet***](https://www.kaggle.com/c/imagenet-object-localization-challenge)\n\n<center> Workflow of LIME Image — How LIME Image actually works?</center>\n<center> Using stratified sampling to improve LIME image explanations. In Proceedings of the AAAI Conference on Artificial Intelligence?</center>\n","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport time\nimport pickle\nimport matplotlib\nimport numpy as np\nimport pandas as pd\nimport scipy.special\nimport json, math,cv2\nimport sys, os, importlib\nimport matplotlib.pyplot as plt\nfrom matplotlib.colors import LinearSegmentedColormap\nfrom skimage.segmentation import mark_boundaries\npd.set_option('display.max_columns', None)\nfrom tqdm.auto import tqdm\n\nmatplotlib.rcParams['text.usetex'] = True\n\n# Stretch Notebook Width to 98% size of the Screen\nfrom IPython.display import display, HTML\ndisplay(HTML(\"<style>.container { width:95% !important; }</style>\"))\n\nimport matplotlib as mpl\nmpl.rcParams.update(mpl.rcParamsDefault)\n\nsys.path.insert(1, '/kaggle/input/xai-with-lime-image-for-computer-vision')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:35:53.2236Z","iopub.execute_input":"2025-04-16T18:35:53.22434Z","iopub.status.idle":"2025-04-16T18:35:54.848207Z","shell.execute_reply.started":"2025-04-16T18:35:53.2243Z","shell.execute_reply":"2025-04-16T18:35:54.847313Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!git clone https://github.com/rashidrao-pk/lime-stratified-examples","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:35:56.933855Z","iopub.execute_input":"2025-04-16T18:35:56.934334Z","iopub.status.idle":"2025-04-16T18:35:58.63303Z","shell.execute_reply.started":"2025-04-16T18:35:56.934307Z","shell.execute_reply":"2025-04-16T18:35:58.631968Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import utils as ut\nsys.path.insert(1, '/kaggle/working/lime-stratified-examples/lime_stratified')\nsys.path.insert(1, '/kaggle/working/lime-stratified-examples')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:36:02.208637Z","iopub.execute_input":"2025-04-16T18:36:02.20898Z","iopub.status.idle":"2025-04-16T18:36:18.342813Z","shell.execute_reply.started":"2025-04-16T18:36:02.208952Z","shell.execute_reply":"2025-04-16T18:36:18.341917Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from dataclasses import dataclass\n@dataclass\nclass Parameters:\n    dummy                : bool  = False\n    model_name           : str   = 'ResNet50'\n    target_seg_no        : int   = 50\n    random_seed          : int = 1234\n    use_stratification   : bool = True\n    top_labels           : int  = 3\n    hide_color           : str = None\n    num_samples          : int = 50\n##########################################################\nparams = ut.Parameters()\nparams.random_state = params.random_seed\n##########################################################\n@dataclass\nclass Paths:\n    dummy                : bool  = False\npaths = Paths()\n\nfrom types import SimpleNamespace\nresults = SimpleNamespace()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:36:47.856856Z","iopub.execute_input":"2025-04-16T18:36:47.857183Z","iopub.status.idle":"2025-04-16T18:36:47.864982Z","shell.execute_reply.started":"2025-04-16T18:36:47.85716Z","shell.execute_reply":"2025-04-16T18:36:47.864013Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"paths.main_results = os.getcwd()\npaths.local_data      =  '/kaggle/input/xai-with-lime-image-for-computer-vision'\npaths.json_file       =  os.path.join(paths.local_data,'data/imagenet_class_index.json')\npaths.imagenet_path   = '/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/'\n\npaths.DS_path         =   os.path.join(paths.imagenet_path, \"test\")\npaths.DS_path_subset         =   os.path.join(paths.local_data, \"data\")\npaths.result_folder   =   os.path.join(paths.main_results, \"result\")\npaths.paper_figures   =   os.path.join(paths.main_results,\"Paper_Figures\")\n\nprint(paths.json_file, os.path.exists(paths.json_file))\nut.check_folders(paths.DS_path)\nut.check_folders(paths.result_folder)\nut.check_folders(paths.paper_figures)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:36:49.876718Z","iopub.execute_input":"2025-04-16T18:36:49.877043Z","iopub.status.idle":"2025-04-16T18:36:49.889397Z","shell.execute_reply.started":"2025-04-16T18:36:49.877017Z","shell.execute_reply":"2025-04-16T18:36:49.888488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"im_ext = ['.jpg', '.JPG', '.jpeg', '.JPEG', '.png', '.PNG']\n\n#get images\ndef get_im(IM_DIR,file_path='imagenet_testset_files.pkl'):\n    if os.path.exists(file_path):\n        with open(file_path, 'rb') as f:\n            im_list = pickle.load(f)\n        return im_list\n    im_list = []\n    i = 1\n    for root, directories, files in tqdm(os.walk(IM_DIR)):\n        for file in files:\n            if any(ext in file for ext in im_ext):\n                im_list.append(os.path.join(root, file))\n            i += 1\n    with open(file_path, 'wb') as f:\n        pickle.dump(im_list, f)\n    return im_list","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:43:18.319978Z","iopub.execute_input":"2025-04-16T18:43:18.320897Z","iopub.status.idle":"2025-04-16T18:43:18.326587Z","shell.execute_reply.started":"2025-04-16T18:43:18.320865Z","shell.execute_reply":"2025-04-16T18:43:18.325715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"im_paths = sorted(get_im(paths.DS_path, file_path = os.path.join(paths.main_results, 'imagenet_testset_files.pkl')))\nim_names = [f.split('/')[-1] for f in im_paths]\n\nim_total = len(im_paths)\nprint(f'total number of images = {im_total}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:43:23.610231Z","iopub.execute_input":"2025-04-16T18:43:23.610847Z","iopub.status.idle":"2025-04-16T18:43:23.759962Z","shell.execute_reply.started":"2025-04-16T18:43:23.610821Z","shell.execute_reply":"2025-04-16T18:43:23.759031Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load Black-box model -> Resnet-50","metadata":{}},{"cell_type":"code","source":"# load pre-trained model and data\nweights_path = \"weights/resnet50_weights.h5\"\nos.makedirs('weights', exist_ok=True)\nmodel = ut.load_model(model_name=params.model_name, \n                      weights_path = weights_path\n                     )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:43:29.133787Z","iopub.execute_input":"2025-04-16T18:43:29.134139Z","iopub.status.idle":"2025-04-16T18:43:32.12378Z","shell.execute_reply.started":"2025-04-16T18:43:29.134117Z","shell.execute_reply":"2025-04-16T18:43:32.122944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# getting ImageNet class names\nclass_names = ut.get_ImageNet_ClassLabels(paths.json_file)\nprint('classes count :',len(class_names))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:43:33.884894Z","iopub.execute_input":"2025-04-16T18:43:33.885201Z","iopub.status.idle":"2025-04-16T18:43:33.896349Z","shell.execute_reply.started":"2025-04-16T18:43:33.885179Z","shell.execute_reply":"2025-04-16T18:43:33.895493Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# data_type = 'local'\ndata_type = 'imagenet'\n########################################################################\nif data_type == 'local':\n    params.image_name = 'bird5.png'\n    results.file = os.path.join(paths.DS_path_subset,params.image_name)\nelif data_type == 'imagenet':\n    params.image_name = 'ILSVRC2012_test_00000125.JPEG'\n    results.file = os.path.join(paths.DS_path,params.image_name)\nparams.image_base_name = params.image_name.split('.')[0]\n########################################################################\nresults.image_to_explain = ut.read_process_image(results.file,model)\nprint(f'{\"Image: \":<15} {results.file} \\n{\"Shape\":<15} {results.image_to_explain.shape}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:43:36.227341Z","iopub.execute_input":"2025-04-16T18:43:36.228034Z","iopub.status.idle":"2025-04-16T18:43:36.27033Z","shell.execute_reply.started":"2025-04-16T18:43:36.228003Z","shell.execute_reply":"2025-04-16T18:43:36.269418Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# for op in sys.path:\n#     # if op in 'working':\n#         print(op)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:43:41.935217Z","iopub.execute_input":"2025-04-16T18:43:41.93556Z","iopub.status.idle":"2025-04-16T18:43:41.939564Z","shell.execute_reply.started":"2025-04-16T18:43:41.935538Z","shell.execute_reply":"2025-04-16T18:43:41.938617Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from lime_stratified.lime import lime_image\nlime_explainer = lime_image.LimeImageExplainer(random_state=params.random_seed)\nfrom lime_stratified.lime.wrappers.scikit_image import SegmentationAlgorithm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:43:43.786851Z","iopub.execute_input":"2025-04-16T18:43:43.787657Z","iopub.status.idle":"2025-04-16T18:43:44.303423Z","shell.execute_reply.started":"2025-04-16T18:43:43.787631Z","shell.execute_reply":"2025-04-16T18:43:44.302442Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Functions for XAI","metadata":{}},{"cell_type":"code","source":"def get_sep():\n    print('-'*100)\ndef plot_segments(results,paths,params, save_plot=True):\n    # num_segments = len(np.unique(results.segments))\n    fig,axes = plt.subplots(1,2, figsize=(6,3))\n    axes[0].imshow(results.image_to_explain); axes[0].set_xticks([]); axes[0].set_yticks([]);  \n    axes[1].imshow(mark_boundaries(results.image_to_explain, results.segments))\n    axes[1].set_xticks([]); axes[1].set_yticks([]); \n    plt.suptitle(f'{results.num_segments} segments')\n    if save_plot:\n        img_name = params.image_name.split('.')[0]\n        # plt.savefig(f'{paths.paper_figures}/{image_name}_image_{num_segments}_segments.pdf', dpi=150, bbox_inches='tight', pad_inches=0.02)\n        plt.savefig(f'{paths.paper_figures}/{img_name}_image_{results.num_segments}_segs.png', transparent=True,dpi=150, bbox_inches='tight', pad_inches=0.02)\n        plt.show()\n######################################################################\ndef compare_lime(results_baseline,results_stratified, params,paths, positive_only=True,\n                      num_features=1000, min_weight_fact=2, cmap='bwr',save_plots=True, verbose=True):\n    v_baseline   = np.max(np.abs(results_baseline.heatmap))\n    v_stratified = np.max(np.abs(results_stratified.heatmap))\n    if verbose:\n        print('v_baseline ', v_baseline)\n        print('v_stratified ', v_stratified)\n        \n    original_shape = results_baseline.image_to_explain.shape[:2]\n    fig, axes = plt.subplots(1, 8, figsize=(16, 4), constrained_layout=True)\n    \n    axes[0].imshow(results_baseline.image_to_explain)\n    axes[1].imshow(mark_boundaries(results_baseline.image_to_explain, results_baseline.segments))\n    ##################################################################\n    temp_base, mask_base = results_baseline.explanation.get_image_and_mask(results_baseline.explanation.top_labels[0],\n                                             positive_only=positive_only,\n                                             num_features=num_features,\n                                             hide_rest=False,\n                                             min_weight=v_baseline / min_weight_fact)\n    \n    # Plot heatmap with matching extent to make it visually same-sized\n    im2 = axes[2].imshow(results_baseline.heatmap, cmap=cmap, vmin=-v_baseline, vmax=v_baseline,\n                         extent=[0, original_shape[1], original_shape[0], 0])\n    fig.colorbar(im2, ax=axes[2], fraction=0.05, pad=0.04)\n    \n    axes[3].imshow(mark_boundaries(temp_base.astype(np.uint8), mask_base))\n    #############################################################################\n    temp_st, mask_st = results_stratified.explanation.get_image_and_mask(results_stratified.explanation.top_labels[0],\n                                             positive_only=positive_only,\n                                             num_features=num_features,\n                                             hide_rest=False,\n                                             min_weight=v_stratified / min_weight_fact)\n    \n    # Plot heatmap with matching extent to make it visually same-sized\n    plt.gca().set_aspect('equal')\n    ut.plot_classification_score(axes[4], results_baseline.explanation,\n                                 results_baseline.X, results_baseline.Y, params.f_x,\n                                plot_everything = False)\n    ###################################################################################################\n    im3 = axes[5].imshow(results_stratified.heatmap, cmap=cmap, vmin=-v_stratified, vmax=v_stratified,\n                         extent=[0, original_shape[1], original_shape[0], 0])\n    fig.colorbar(im3, ax=axes[5], fraction=0.05, pad=0.04)\n    ###############\n    axes[6].imshow(mark_boundaries(temp_st.astype(np.uint8), mask_st))\n    ##################################################################################################\n    plt.gca().set_aspect('equal')\n    ut.plot_classification_score(axes[7], results_stratified.explanation,\n                                 results_stratified.X, results_stratified.Y, params.f_x,\n                                plot_everything = False)\n    \n    # Uniform look for all axes\n    for ax in axes:\n        ax.set_xticks([])\n        ax.set_yticks([])\n        ax.set_aspect('equal', adjustable='box')\n        # fr\"{method} $\\mathbf{{CV}}$\"\n    rcy_ttl_bl = f'$RC-LIME = \\\\mathbf{{ {ut.get_RCY(results_baseline.Y, params.f_x):.3} }}$'\n    rcy_ttl_sl = f'$RC-St-LIME = \\\\mathbf{{ {ut.get_RCY(results_stratified.Y, params.f_x):.3} }}$'\n    \n    cvb_bl = fr'LIME $\\mathbf{{CV}}$ : {results_baseline.cv_beta:0.4}'\n    cvb_sl = fr'St-LIME $\\mathbf{{CV}}$ : {results_stratified.cv_beta:0.4}'\n    \n    ttl_ls = ['input', f'Segments:{results_baseline.num_segments}',cvb_bl,'LIME-feat',rcy_ttl_bl,cvb_sl ,'St-LIME-feat',rcy_ttl_sl]\n    \n    for ax,ttl in zip(axes, ttl_ls):\n        ax.set_title(ttl)\n    \n    plt.suptitle(f'predicted as {class_names[params.predicted_cls_idx]}'\n                 f' f(x)={params.f_x:.5}  g(x1)={results_baseline.g_x:.5}, g(x1)={results_stratified.g_x:.5}',\n                fontsize=16)\n    # plt.tight_layout()\n    # plt.subplots_adjust()\n    if save_plots:\n        img_name = params.image_name.split('.')[0]\n    \n    plt.show()\n\n\n    \ndef plot_explanations(results, params,paths , positive_only=True,save_plots=False,\n                      num_features=1000, min_weight_fact=2, cmap='bwr', verbose=True):\n    image_name = params.image_name.split('.')[0]\n    v = np.max(np.abs(results.heatmap))\n    if verbose:\n        print(f'{\"max imp \":<15} =  {v}')\n    \n    temp_1, mask_1 = results.explanation.get_image_and_mask(results.explanation.top_labels[0],\n                                             positive_only=positive_only,\n                                             num_features=num_features,\n                                             hide_rest=True,\n                                             min_weight=v / min_weight_fact)\n\n    temp_2, mask_2 = results.explanation.get_image_and_mask(results.explanation.top_labels[0],\n                                             positive_only=positive_only,\n                                             num_features=num_features,\n                                             hide_rest=False,\n                                             min_weight=v / min_weight_fact)\n\n    # Resize heatmap to match the original image shape (height, width)\n    original_shape = temp_1.shape[:2]\n    heatmap = results.heatmap\n    if heatmap.shape[:2] != original_shape:\n        heatmap = cv2.resize(heatmap, (original_shape[1], original_shape[0]), interpolation=cv2.INTER_NEAREST)\n\n    # Create subplots with equal aspect and no padding\n    fig, axes = plt.subplots(1, 5, figsize=(12, 4), constrained_layout=True)\n    axes[0].imshow(results.image_to_explain)\n    axes[1].imshow(mark_boundaries(results.image_to_explain, results.segments))\n    \n    # axes[2].imshow(mark_boundaries(temp_1.astype(np.uint8), mask_1))\n\n    # Classification Score Plot\n    # fig, ax = plt.subplots(1, 1, figsize=(2.5, 2.5))\n    \n    # plt.title()\n    \n    axes[2].imshow(mark_boundaries(temp_2.astype(np.uint8), mask_2))\n    ##################################################################################\n    plt.gca().set_aspect('equal')\n    ut.plot_classification_score(axes[3], results.explanation, results.X, results.Y, params.f_x)\n    ##################################################################################\n    # Plot heatmap with matching extent to make it visually same-sized\n    im = axes[4].imshow(heatmap, cmap=cmap, vmin=-v, vmax=v, extent=[0, original_shape[1], original_shape[0], 0])\n    fig.colorbar(im, ax=axes[4], fraction=0.05, pad=0.04)\n    \n    # Uniform look for all axes\n    for ax in axes:\n        ax.set_xticks([])\n        ax.set_yticks([])\n        ax.set_aspect('equal', adjustable='box')\n    rcy_ttl = f'$RC(Y) = \\\\mathbf{{ {ut.get_RCY(results.Y, params.f_x):.3} }}$'\n    ttl_ls = ['input', 'segments','feat',rcy_ttl, f'Beta : {results.cv_beta:0.4}']\n    for ax,ttl in zip(axes, ttl_ls):\n        ax.set_title(ttl)\n    \n    plt.suptitle(f'predicted as {class_names[params.predicted_cls_idx]}  '\n                      f'f(x)={params.f_x:.5}  g(x)={results.g_x:.5}')\n    # plt.tight_layout(w_pad=0.05, h_pad=0.05)\n    # plt.subplots_adjust(hspace=0.05, wspace=0.05)\n    if save_plots:\n        # Save files\n        for fmt in ['svg', 'pdf', 'png']:\n            plt.savefig(f'{paths.paper_figures}/{image_name}_image_mask_heatmap_single.{fmt}',\n                    dpi=150, bbox_inches='tight', pad_inches=0.02,\n                    transparent=(fmt == 'png'))\n    plt.show()\n#######################################################################\n\ndef get_beta_exp(results, verbose=False):\n    xpld_cls = results.explanation.top_labels[0]\n    results.g_x = results.explanation.local_pred[xpld_cls][0]\n    results.beta = get_beta_from_expl(explanation=results.explanation)    \n    results.std_beta = np.std((results.beta))\n    results.mean_beta = np.mean((results.beta))\n    \n    if verbose:\n        print('g(x) \\t\\t= ', results.g_x)\n        print('sum(beta) \\t= ', np.sum(results.beta))\n        print('CV(beta) \\t= ',results.std_beta/results.mean_beta)\n    return results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T19:26:19.230979Z","iopub.execute_input":"2025-04-16T19:26:19.231735Z","iopub.status.idle":"2025-04-16T19:26:19.258131Z","shell.execute_reply.started":"2025-04-16T19:26:19.231707Z","shell.execute_reply":"2025-04-16T19:26:19.257189Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Black Box Model","metadata":{}},{"cell_type":"code","source":"from tensorflow.keras.applications.resnet50 import preprocess_input\ndef bb_predict(imgs):\n    # On some platform, you will need model.predict(..) instead of model(..)\n    return model.predict(preprocess_input(imgs), verbose=False)\n#     return model(preprocess_input(imgs))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:43:50.269936Z","iopub.execute_input":"2025-04-16T18:43:50.270287Z","iopub.status.idle":"2025-04-16T18:43:50.274813Z","shell.execute_reply.started":"2025-04-16T18:43:50.270233Z","shell.execute_reply":"2025-04-16T18:43:50.274015Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predicted = bb_predict(np.array([results.image_to_explain]))\n(params.predicted_cls_idx,params.f_x,\\\n params.predicted_cls_lbl) =  ut.get_class_idx_label_score (predicted,class_names)\nprint(params.predicted_cls_idx, params.predicted_cls_lbl, params.f_x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:43:51.971096Z","iopub.execute_input":"2025-04-16T18:43:51.971431Z","iopub.status.idle":"2025-04-16T18:43:54.956145Z","shell.execute_reply.started":"2025-04-16T18:43:51.971407Z","shell.execute_reply":"2025-04-16T18:43:54.955108Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create Segments in Image","metadata":{}},{"cell_type":"markdown","source":"### Find Max Dist for Segmentation algorithm based on Target number of Segments","metadata":{}},{"cell_type":"code","source":"csv_segments_file = os.path.join(paths.result_folder,'segments_db.csv')\nparams.max_dist = ut.get_max_dist_load(params,results,csv_segments_file, verbose=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:45:02.863871Z","iopub.execute_input":"2025-04-16T18:45:02.864247Z","iopub.status.idle":"2025-04-16T18:45:02.875859Z","shell.execute_reply.started":"2025-04-16T18:45:02.864224Z","shell.execute_reply":"2025-04-16T18:45:02.874822Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Create Segments using Hyper searched parameters","metadata":{}},{"cell_type":"code","source":"results.segments,results.num_segments,segmenter_fn = ut.own_seg(results.image_to_explain,\n                                                                md=params.max_dist,ks=4,\n                                                                random_seed=params.random_seed,ratio=0.2)\nprint(f'num_segments created --> {results.num_segments} - {segmenter_fn}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:45:07.450536Z","iopub.execute_input":"2025-04-16T18:45:07.450988Z","iopub.status.idle":"2025-04-16T18:45:08.155361Z","shell.execute_reply.started":"2025-04-16T18:45:07.450956Z","shell.execute_reply":"2025-04-16T18:45:08.154423Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"explanation = lime_explainer.explain_instance(results.image_to_explain,             # image being explained\n              bb_predict,                   # prediction model \n              labels=class_names,           # classes names from ImageNet dataset\n              segmentation_fn=segmenter_fn, # custom Segmenter function to generate exactly same superpixels\n              top_labels=params.top_labels,       # top explanation\n              hide_color=params.hide_color,       # superpixel replacement strategy\n              use_stratification=params.use_stratification, # Boolean value to switch the proposed method to be used or not \n              # batch_size=100,               # batch size\n              num_samples=params.num_samples)             # no of \nresults.X, results.all_Ys, results.explanation = explanation\nresults.Y = results.all_Ys[:, params.predicted_cls_idx]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:53:20.057844Z","iopub.execute_input":"2025-04-16T18:53:20.058184Z","iopub.status.idle":"2025-04-16T18:53:29.540965Z","shell.execute_reply.started":"2025-04-16T18:53:20.058161Z","shell.execute_reply":"2025-04-16T18:53:29.540319Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results = ut.get_beta_exp(results)\nresults.heatmap = ut.heatmap_from_beta(segments=results.segments, beta=results.beta)\nresults.cv_beta = ut.get_CV_beta(results.beta)\nplot_explanations(results,params,paths,positive_only=False,\n                     num_features=100,min_weight_fact=3,cmap='bwr', verbose=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T19:03:20.42346Z","iopub.execute_input":"2025-04-16T19:03:20.423784Z","iopub.status.idle":"2025-04-16T19:03:23.062604Z","shell.execute_reply.started":"2025-04-16T19:03:20.42376Z","shell.execute_reply":"2025-04-16T19:03:23.061538Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f'{\"local_pred \":<15} -> {results.explanation.local_pred}')\nprint(f'{\"score \":<15} -> {results.explanation.score}')\nprint(f'{\"segments\":<15} -> {len(np.unique(results.explanation.segments))}')\nprint(f'{\"top_labels \":<15} -> {results.explanation.top_labels}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T18:49:48.219116Z","iopub.execute_input":"2025-04-16T18:49:48.219864Z","iopub.status.idle":"2025-04-16T18:49:48.225746Z","shell.execute_reply.started":"2025-04-16T18:49:48.219837Z","shell.execute_reply":"2025-04-16T18:49:48.225043Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"def get_flow(params=None,paths=None, verbose=False,plot_full=True,verbose_seg=True):\n    results = SimpleNamespace()\n    \n    results.file = os.path.join(paths.DS_path,params.image_name)\n    params.image_base_name = params.image_name.split('.')[0]\n    results.image_to_explain = ut.read_process_image(results.file,model)\n    if verbose:\n        print(f'{params.image_name} loaded with shape : {results.image_to_explain.shape}')\n        get_sep()\n    ##########################################################################################################################\n    predicted = bb_predict(np.array([results.image_to_explain]))\n    \n    (params.predicted_cls_idx,params.f_x,params.predicted_cls_lbl) =  ut.get_class_idx_label_score (predicted,class_names)\n\n    # predicted_cls = np.argmax(predicted[0])\n    # f_x = predicted[0][predicted_cls]\n    if verbose:\n        print('Predicted Class\\t\\t: \\t',params.predicted_cls_lbl,\n              '\\nClass Probability\\t:\\t', params.f_x,\n              '\\nPredicted Class Index\\t:\\t', params.predicted_cls_idx)\n        get_sep()\n    ##########################################################################################################################\n    params.max_dist = ut.get_max_dist_load(params,results,csv_segments_file, verbose=verbose_seg)\n    # max_dist,_,_,_ = ut.search_segment_number(results.image_to_explain, target_seg_no=params.target_seg_no)\n    if verbose:\n        print(f'{params.target_seg_no} segments requires : max_dist: {params.max_dist}')\n    results.segments,results.num_segments,segmenter_fn = ut.own_seg(results.image_to_explain,\n                                                                    md=params.max_dist,\n                                                                    ks=4,\n                                                                    random_seed=params.random_seed,\n                                                                    ratio=0.2)\n    ##########################################################################################################################\n    lime_explainer = lime_image.LimeImageExplainer(random_state=params.random_seed) \n    # Boolean value to switch the proposed method to be used or not\n    if verbose:\n        print('TOP Labels: ', params.top_labels)\n        print('hide_color: ', params.hide_color)\n        print('use_stratification: ', params.use_stratification)\n        print('num_samples: ', params.num_samples)\n    \n    explanation = lime_explainer.explain_instance(results.image_to_explain,            # image being explained\n                                          bb_predict,                   # prediction model \n                                          labels=class_names,           # classes names from ImageNet dataset\n                                          segmentation_fn=segmenter_fn, # custom Segmenter function to generate exactly same superpixels\n                                          top_labels=params.top_labels,                 # top explanation\n                                          hide_color=params.hide_color,              # superpixel replacement strategy\n                                          use_stratification=params.use_stratification, # Boolean value to switch the proposed method to be used or not \n                                          num_samples=params.num_samples)             # no of samples\n    results.X, results.all_Ys, results.explanation = explanation\n    results.Y = results.all_Ys[:, params.predicted_cls_idx]\n\n    if verbose:\n        print(explanation.top_labels[0])\n    results = ut.get_beta_exp(results)\n    results.heatmap = ut.heatmap_from_beta(segments=results.segments, beta=results.beta)\n    results.cv_beta = ut.get_CV_beta(results.beta)\n    if verbose:\n        print('CV Value ', ut.get_CV_beta(results_baseline.beta))\n    if plot_full:\n        plot_explanations(results,params,paths,positive_only=params.positive_only,\n                      num_features=params.num_features,\n                      min_weight_fact=params.min_weight_fact,\n                      cmap=params.cmap,\n                         verbose=verbose)\n    \n    return results\ndef plot_rc_score(results,params):\n    \n    # Classification Score Plot\n    fig, ax = plt.subplots(1, 1, figsize=(2.5, 2.5))\n    plt.gca().set_aspect('equal')\n    ut.plot_classification_score(ax, results.explanation, results.X, results.Y, params.f_x)\n    plt.title(f'$RC(Y) = \\\\mathbf{{ {ut.get_RCY(results.Y, params.f_x):.3} }}$')\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T19:05:43.550835Z","iopub.execute_input":"2025-04-16T19:05:43.551217Z","iopub.status.idle":"2025-04-16T19:05:43.56477Z","shell.execute_reply.started":"2025-04-16T19:05:43.551181Z","shell.execute_reply":"2025-04-16T19:05:43.563889Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### HyperParameters for No of Segments \n- target_seg_no = 100\n- \n### Feature Importance Computing\n- **Budget** (num_samples) is Budget for Explanation Generation \n#################################\n### Feature Visualization\n- num_features = 100\n- positive_only = False   # Parameter for Feature Importance Visualization\n- min_weight_fact = 2     # Parameter for Feature Importance Visualization\n- colormap = 'bwr'","metadata":{}},{"cell_type":"code","source":"params.target_seg_no = 100   # Segments Generation\n#################################\nparams.num_samples = 100   # Budget for Explanation Generation \n#################################\nparams.num_features = 100\nparams.positive_only = False   # Parameter for Feature Importance Visualization\nparams.min_weight_fact = 2     # Parameter for Feature Importance Visualization\nparams.cmap = 'bwr'\n\nparams.image_name = 'ILSVRC2012_test_00000125.JPEG'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T19:03:41.584158Z","iopub.execute_input":"2025-04-16T19:03:41.584516Z","iopub.status.idle":"2025-04-16T19:03:41.589872Z","shell.execute_reply.started":"2025-04-16T19:03:41.58449Z","shell.execute_reply":"2025-04-16T19:03:41.588876Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## LIME Without Stratified Sampling","metadata":{}},{"cell_type":"code","source":"params.use_stratification = False\nresults_baseline = get_flow(params=params,paths=paths, verbose_seg = False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T19:06:07.342516Z","iopub.execute_input":"2025-04-16T19:06:07.342867Z","iopub.status.idle":"2025-04-16T19:06:21.694483Z","shell.execute_reply.started":"2025-04-16T19:06:07.34284Z","shell.execute_reply":"2025-04-16T19:06:21.693303Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## LIME with Stratified Sampling","metadata":{}},{"cell_type":"code","source":"params.use_stratification = True\nresults_stratified = get_flow(params=params,paths=paths, verbose_seg = False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T19:06:26.080241Z","iopub.execute_input":"2025-04-16T19:06:26.080597Z","iopub.status.idle":"2025-04-16T19:06:41.298794Z","shell.execute_reply.started":"2025-04-16T19:06:26.080572Z","shell.execute_reply":"2025-04-16T19:06:41.297833Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Compare both Feature Importances\n","metadata":{}},{"cell_type":"code","source":"compare_lime(results_baseline,results_stratified,\n             params,paths,\n             positive_only=False,\n             num_features=20,\n             min_weight_fact=2, \n             cmap='bwr', #\n             verbose=True\n            )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T19:26:22.625874Z","iopub.execute_input":"2025-04-16T19:26:22.62619Z","iopub.status.idle":"2025-04-16T19:26:23.403766Z","shell.execute_reply.started":"2025-04-16T19:26:22.626151Z","shell.execute_reply":"2025-04-16T19:26:23.402715Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"params.target_seg_no = 100   # Segments Generation\n#################################\nparams.num_samples = 1000   # Budget for Explanation Generation \n#################################\nparams.num_features = 100\nparams.positive_only = False   # Parameter for Feature Importance Visualization\nparams.min_weight_fact = 2     # Parameter for Feature Importance Visualization\nparams.cmap = 'bwr'\n\nparams.image_name = 'ILSVRC2012_test_00000125.JPEG'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T19:27:00.315425Z","iopub.execute_input":"2025-04-16T19:27:00.31573Z","iopub.status.idle":"2025-04-16T19:27:00.321337Z","shell.execute_reply.started":"2025-04-16T19:27:00.315707Z","shell.execute_reply":"2025-04-16T19:27:00.320485Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### RUN ON MULTIPLE IMAGES FROM TEST SET OF IMAGENET ","metadata":{}},{"cell_type":"code","source":"# TEST_ON_MULTI_IMAGE = True\nTEST_ON_MULTI_IMAGE = False\n\nif TEST_ON_MULTI_IMAGE:\n    selected_images = []\n    image_no = [\n                114,     147,      60,       144,  66\n                ]\n    for im_idx,ino in enumerate(image_no):\n        selected_images.append(f'ILSVRC2012_test_{ino:08}.JPEG')\n        \n    print(selected_images)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T19:27:53.796343Z","iopub.execute_input":"2025-04-16T19:27:53.796647Z","iopub.status.idle":"2025-04-16T19:27:53.801484Z","shell.execute_reply.started":"2025-04-16T19:27:53.796624Z","shell.execute_reply":"2025-04-16T19:27:53.800621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if TEST_ON_MULTI_IMAGE:\n    for imn in selected_images:\n        results = SimpleNamespace()\n        params.image_name = imn\n        ############################################################################################\n        params.use_stratification = False\n        results_baseline = get_flow(params=params,paths=paths, plot_full=False, verbose_seg=False)\n        ############################################################################################\n        params.use_stratification = True\n        results_stratified = get_flow(params=params,paths=paths, plot_full=False, verbose_seg=False)\n        ############################################################################################\n        compare_lime(results_baseline,results_stratified,\n                     params,paths,\n                     positive_only=False,\n                     num_features=20,\n                     min_weight_fact=2, \n                     cmap='bwr', #\n                     verbose=False\n                    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-16T19:28:01.857108Z","iopub.execute_input":"2025-04-16T19:28:01.857925Z","iopub.status.idle":"2025-04-16T19:28:01.863138Z","shell.execute_reply.started":"2025-04-16T19:28:01.85788Z","shell.execute_reply":"2025-04-16T19:28:01.862147Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# END","metadata":{}}]}