DYNAMAXD/ExplainabilityOnModel
010
1import os2import base643import numpy as np4import cv25import tensorflow as tf6from fastapi import FastAPI, UploadFile, File7from fastapi.responses import JSONResponse8from fastapi.staticfiles import StaticFiles9from fastapi.middleware.cors import CORSMiddleware10import io11 12app = FastAPI(title="Brain Tumor XAI API")13 14app.add_middleware(15 CORSMiddleware,16 allow_origins=["*"],17 allow_methods=["*"],18 allow_headers=["*"],19)20 21# Load the model globally22model_path = "resnet_finetuned_model_v1.keras"23if os.path.exists(model_path):24 print("Loading model...")25 model = tf.keras.models.load_model(model_path)26 print("Model loaded successfully.")27else:28 print(f"Warning: Model not found at {model_path}.")29 model = None30 31class_names = ['glioma', 'meningioma', 'notumor', 'pituitary']32 33def img_to_base64(img_array):34 img_bgr = cv2.cvtColor(img_array, cv2.COLOR_RGB2BGR)35 _, buffer = cv2.imencode('.jpg', img_bgr)36 return base64.b64encode(buffer).decode('utf-8')37 38def superimpose_heatmap(original_img, heatmap_arr):39 # Resize heatmap to match original image if necessary, but original_img here is resized? 40 # original_img usually is kept its own size, but for UI, we will resize both to e.g. 256x256 or just the original size41 heatmap_resized = cv2.resize(heatmap_arr, (original_img.shape[1], original_img.shape[0]))42 # Convert heatmap to RGB 43 heatmap_colored = cv2.applyColorMap(np.uint8(255 * heatmap_resized), cv2.COLORMAP_JET)44 heatmap_rgb = cv2.cvtColor(heatmap_colored, cv2.COLOR_BGR2RGB)45 46 superimposed_img = heatmap_rgb * 0.5 + original_img * 0.547 return np.uint8(superimposed_img)48 49def compute_gradcam(img_array, model, layer_name="conv5_block3_out"):50 grad_model = tf.keras.models.Model(51 [model.inputs],52 [model.get_layer(layer_name).output, model.output]53 )54 img_tensor = tf.convert_to_tensor(img_array, dtype=tf.float32)55 with tf.GradientTape() as tape:56 conv_outputs, predictions = grad_model(img_tensor)57 class_idx = tf.argmax(predictions[0])58 loss = predictions[:, class_idx]59 grads = tape.gradient(loss, conv_outputs)60 pooled_grads = tf.reduce_mean(grads, axis=(0,1,2))61 conv_outputs = conv_outputs[0]62 heatmap = conv_outputs @ pooled_grads[..., tf.newaxis]63 heatmap = tf.squeeze(heatmap)64 heatmap = tf.maximum(heatmap, 0)65 heatmap /= tf.reduce_max(heatmap) + 1e-866 return heatmap.numpy()67 68def compute_gradcam_plus_plus(img_array, model, layer_name="conv5_block3_out"):69 grad_model = tf.keras.models.Model(70 [model.inputs],71 [model.get_layer(layer_name).output, model.output]72 )73 img_tensor = tf.convert_to_tensor(img_array, dtype=tf.float32)74 with tf.GradientTape() as tape:75 conv_outputs, predictions = grad_model(img_tensor)76 class_idx = tf.argmax(predictions[0])77 loss = predictions[:, class_idx]78 grads = tape.gradient(loss, conv_outputs)79 conv_outputs = conv_outputs[0]80 grads = grads[0]81 grads_power_2 = grads ** 282 grads_power_3 = grads ** 383 sum_activations = tf.reduce_sum(conv_outputs, axis=(0,1))84 alpha_num = grads_power_285 alpha_denom = 2*grads_power_2 + grads_power_3 * sum_activations86 alpha_denom = tf.where(alpha_denom != 0.0, alpha_denom, tf.ones_like(alpha_denom))87 alphas = alpha_num / alpha_denom88 weights = tf.reduce_sum(alphas * tf.nn.relu(grads), axis=(0,1))89 heatmap = tf.reduce_sum(weights * conv_outputs, axis=-1)90 heatmap = tf.maximum(heatmap,0)91 heatmap /= tf.reduce_max(heatmap) + 1e-892 return heatmap.numpy()93 94def compute_scorecam(img_array, model, layer_name="conv5_block3_out"):95 grad_model = tf.keras.models.Model(96 [model.inputs],97 [model.get_layer(layer_name).output]98 )99 conv_outputs = grad_model(img_array)100 conv_outputs = conv_outputs[0].numpy()101 heatmap = np.zeros(conv_outputs.shape[:2])102 103 # Restrict channel iteration to speed up server response if it's too slow:104 # We will use sub-sampling of channels to speed it up to ~50 channels105 num_channels = min(50, conv_outputs.shape[-1])106 channels_to_process = np.linspace(0, conv_outputs.shape[-1]-1, num_channels, dtype=int)107 108 for i in channels_to_process:109 activation = conv_outputs[:,:,i]110 activation_norm = (activation - activation.min()) / (activation.max() - activation.min() + 1e-8)111 activation_resized = cv2.resize(activation_norm, (224,224))112 masked_img = img_array * activation_resized[...,np.newaxis]113 preds = model.predict(masked_img, verbose=0)114 score = np.max(preds)115 heatmap += activation_norm * score116 117 heatmap = np.maximum(heatmap, 0)118 heatmap /= np.max(heatmap) + 1e-8119 return heatmap120 121def compute_occlusion(img_array, model, patch_size=20, stride=20):122 img = img_array.copy()123 h = img.shape[1]124 w = img.shape[2]125 heatmap = np.zeros((h, w))126 counts = np.zeros((h, w))127 128 preds = model.predict(img_array, verbose=0)129 class_idx = np.argmax(preds)130 baseline_score = preds[0][class_idx]131 132 # Stride of 20 to speed up API response133 for y in range(0, h - patch_size + 1, stride):134 for x in range(0, w - patch_size + 1, stride):135 occluded = img.copy()136 occluded[:, y:y+patch_size, x:x+patch_size, :] = 0137 preds_occ = model.predict(occluded, verbose=0)138 score = preds_occ[0][class_idx]139 importance = baseline_score - score140 heatmap[y:y+patch_size, x:x+patch_size] += importance141 counts[y:y+patch_size, x:x+patch_size] += 1142 143 heatmap = heatmap / (counts + 1e-8)144 heatmap = np.maximum(heatmap, 0)145 heatmap = heatmap / (np.max(heatmap) + 1e-8)146 heatmap = cv2.GaussianBlur(heatmap, (11,11), 0)147 return heatmap148 149def compute_eigencam(img_array, model, layer_name="conv5_block3_out"):150 feature_model = tf.keras.models.Model(151 [model.inputs],152 [model.get_layer(layer_name).output]153 )154 feature_maps = feature_model.predict(img_array, verbose=0)155 feature_maps = feature_maps[0]156 h, w, c = feature_maps.shape157 reshaped = feature_maps.reshape((h*w, c))158 reshaped = reshaped - np.mean(reshaped, axis=0)159 U, S, Vt = np.linalg.svd(reshaped, full_matrices=False)160 principal_component = Vt[0]161 heatmap = np.dot(reshaped, principal_component)162 heatmap = heatmap.reshape(h, w)163 heatmap = np.maximum(heatmap, 0)164 heatmap = heatmap / (np.max(heatmap) + 1e-8)165 return heatmap166 167def compute_layercam(img_array, model, layer_name="conv5_block3_out"):168 grad_model = tf.keras.models.Model(169 [model.inputs],170 [model.get_layer(layer_name).output, model.output]171 )172 img_tensor = tf.convert_to_tensor(img_array, dtype=tf.float32)173 with tf.GradientTape() as tape:174 conv_outputs, predictions = grad_model(img_tensor)175 class_idx = tf.argmax(predictions[0])176 loss = predictions[:, class_idx]177 grads = tape.gradient(loss, conv_outputs)178 conv_outputs = conv_outputs[0]179 grads = grads[0]180 positive_grads = tf.nn.relu(grads)181 layercam = positive_grads * conv_outputs182 heatmap = tf.reduce_sum(layercam, axis=-1)183 heatmap = tf.maximum(heatmap,0)184 heatmap /= tf.reduce_max(heatmap) + 1e-8185 return heatmap.numpy()186 187@app.post("/predict")188async def predict(file: UploadFile = File(...)):189 if model is None:190 return JSONResponse(status_code=500, content={"error": "Model not loaded"})191 192 contents = await file.read()193 nparr = np.frombuffer(contents, np.uint8)194 img = cv2.imdecode(nparr, cv2.IMREAD_COLOR)195 original_img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)196 197 # Process for ResNet198 img_resized = cv2.resize(original_img, (224, 224))199 img_array = np.array(img_resized, dtype=np.float32)200 img_array = np.expand_dims(img_array, axis=0)201 img_array = tf.keras.applications.resnet50.preprocess_input(img_array)202 203 preds = model.predict(img_array)204 pred_index = int(np.argmax(preds))205 pred_class = class_names[pred_index]206 confidence = float(preds[0][pred_index])207 208 # Compute CAMs209 # Note: to ensure high performance locally, we limit original size base64 to scaled versions if too large,210 # but 224x224 to 500x500 is fine. Here we scale original_img for visualization to 300x300211 viz_img = cv2.resize(original_img, (300, 300))212 213 orig_b64 = img_to_base64(viz_img)214 215 hm_gradcam = compute_gradcam(img_array, model)216 hm_gradcam_pp = compute_gradcam_plus_plus(img_array, model)217 hm_scorecam = compute_scorecam(img_array, model)218 hm_occlusion = compute_occlusion(img_array, model)219 hm_eigencam = compute_eigencam(img_array, model)220 hm_layercam = compute_layercam(img_array, model)221 222 cams = {223 "gradcam": img_to_base64(superimpose_heatmap(viz_img, hm_gradcam)),224 "gradcam_pp": img_to_base64(superimpose_heatmap(viz_img, hm_gradcam_pp)),225 "scorecam": img_to_base64(superimpose_heatmap(viz_img, hm_scorecam)),226 "occlusion": img_to_base64(superimpose_heatmap(viz_img, hm_occlusion)),227 "eigencam": img_to_base64(superimpose_heatmap(viz_img, hm_eigencam)),228 "layercam": img_to_base64(superimpose_heatmap(viz_img, hm_layercam))229 }230 231 return {232 "pred_class": pred_class,233 "confidence": confidence,234 "original_image": orig_b64,235 "cams": cams236 }237 238print("Mounting static files now...")239os.makedirs("static", exist_ok=True)240app.mount("/", StaticFiles(directory="static", html=True), name="static")241 242 