import io import base64 import time import cv2 import numpy as np import torch import torchvision.transforms as transforms from PIL import Image from pytorch_grad_cam import GradCAM from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget from pytorch_grad_cam.utils.image import show_cam_on_image from model_loader import load_wbc_model, load_skin_model # Class labels WBC_CLASSES = ["Basophil", "Eosinophil", "Erythroblast", "IG", "Lymphocyte", "Monocyte", "Neutrophil", "Platelet"] WBC_CLASSES_TF = ["Basophil", "Eosinophil", "Erythroblast", "IG", "Lymphocyte", "Monocyte", "Neutrophil", "Platelet", "RBC", "WBC", "Other"] SKIN_CLASSES = [ "Benign keratosis-like lesions", "Basal cell carcinoma", "Actinic keratoses", "Vascular lesions", "Melanocytic nevi", "Melanoma", "Dermatofibroma" ] # Image normalization parameters IMAGE_SIZE = 224 NORM_MEAN = [0.485, 0.456, 0.406] NORM_STD = [0.229, 0.224, 0.225] def get_wbc_transforms(): return transforms.Compose([ transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)), transforms.ToTensor(), transforms.Normalize(mean=NORM_MEAN, std=NORM_STD) ]) def generate_wbc_gradcam(model, input_tensor, target_class_idx, raw_image_np): """ Generates a Grad-CAM overlay for the WBC ResNet-18 model. """ try: target_layers = [model[0][7][-1]] cam = GradCAM(model=model, target_layers=target_layers) targets = [ClassifierOutputTarget(target_class_idx)] grayscale_cam = cam(input_tensor=input_tensor, targets=targets)[0, :] rgb_img = cv2.resize(raw_image_np, (IMAGE_SIZE, IMAGE_SIZE)) / 255.0 cam_image = show_cam_on_image(rgb_img, grayscale_cam, use_rgb=True) return cam_image except Exception as e: print(f"Error generating WBC Grad-CAM: {str(e)}") return None def generate_wbc_gradcam_tf(model, input_tensor, target_class_idx, raw_image_np): """ Generates a Grad-CAM overlay for the custom Keras WBC CNN model. """ try: import tensorflow as tf with tf.GradientTape() as tape: x = tf.convert_to_tensor(input_tensor) curr_x = x conv_outputs = None for layer in model.layers: curr_x = layer(curr_x) if isinstance(layer, tf.keras.layers.Conv2D): conv_outputs = curr_x predictions = curr_x loss = predictions[:, target_class_idx] if conv_outputs is None: print("Error: No Conv2D layer found in Keras model.") return None grads = tape.gradient(loss, conv_outputs) pooled_grads = tf.reduce_mean(grads, axis=(0, 1, 2)) conv_outputs_val = conv_outputs[0] heatmap = conv_outputs_val @ pooled_grads[..., tf.newaxis] heatmap = tf.squeeze(heatmap) heatmap = tf.maximum(heatmap, 0.0) max_val = tf.math.reduce_max(heatmap) if max_val > 0: heatmap = heatmap / max_val heatmap = heatmap.numpy() heatmap_resized = cv2.resize(heatmap, (128, 128)) rgb_img = cv2.resize(raw_image_np, (128, 128)) heatmap_color = cv2.applyColorMap(np.uint8(255 * heatmap_resized), cv2.COLORMAP_JET) heatmap_color = cv2.cvtColor(heatmap_color, cv2.COLOR_BGR2RGB) blend_image = cv2.addWeighted(rgb_img, 0.6, heatmap_color, 0.4, 0) return blend_image except Exception as e: print(f"Error generating TensorFlow WBC Grad-CAM: {str(e)}") return None def generate_skin_attention_map(model, inputs, target_class_idx, raw_image_np): """ Generates a self-attention heatmap for the Vision Transformer skin cancer model. """ try: with torch.no_grad(): outputs = model(**inputs, output_attentions=True) attentions = outputs.attentions[-1] mean_attn = attentions.mean(dim=1)[0] cls_attn = mean_attn[0, 1:] grid_size = int(np.sqrt(cls_attn.size(0))) heatmap_grid = cls_attn.reshape(grid_size, grid_size).cpu().numpy() heatmap_grid = (heatmap_grid - heatmap_grid.min()) / (heatmap_grid.max() - heatmap_grid.min() + 1e-8) heatmap_resized = cv2.resize(heatmap_grid, (IMAGE_SIZE, IMAGE_SIZE)) heatmap_color = cv2.applyColorMap(np.uint8(255 * heatmap_resized), cv2.COLORMAP_JET) heatmap_color = cv2.cvtColor(heatmap_color, cv2.COLOR_BGR2RGB) rgb_img = cv2.resize(raw_image_np, (IMAGE_SIZE, IMAGE_SIZE)) blend_image = cv2.addWeighted(rgb_img, 0.6, heatmap_color, 0.4, 0) return blend_image except Exception as e: print(f"Error generating Skin Attention Map: {str(e)}") return None def segment_cells(img_bgr, min_area=500): lab = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2LAB) blur = cv2.GaussianBlur(lab[:,:,1],(7,7),0) _,thresh = cv2.threshold(blur,0,255, cv2.THRESH_BINARY_INV+cv2.THRESH_OTSU) k = cv2.getStructuringElement(cv2.MORPH_ELLIPSE,(9,9)) cl = cv2.morphologyEx(thresh,cv2.MORPH_CLOSE,k,iterations=2) cl = cv2.morphologyEx(cl,cv2.MORPH_OPEN,k,iterations=1) cnts,_ = cv2.findContours(cl,cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) boxes = [] for c in cnts: if cv2.contourArea(c) < min_area: continue x,y,w,h = cv2.boundingRect(c) if 0.3 < w/max(h,1) < 3.0: boxes.append((x,y,w,h)) return sorted(boxes, key=lambda b:b[2]*b[3], reverse=True) def predict_image(image_bytes: bytes, filename: str, module_type: str): """ Performs inference and generates heatmaps. """ start_time = time.time() try: image = Image.open(io.BytesIO(image_bytes)).convert("RGB") except Exception as e: raise ValueError(f"Invalid image content: {str(e)}") raw_image_np = np.array(image) if module_type == "blood_cell": wbc_model_data = load_wbc_model() framework = wbc_model_data["framework"] model = wbc_model_data["model"] if framework == "tensorflow": import tensorflow as tf import matplotlib matplotlib.use('Agg') import matplotlib.pyplot as plt from collections import Counter # Setup class names list in the exact order as Colab training CLASS_NAMES = ['basophil', 'eosinophil', 'erythroblast', 'ig', 'Lymphocyte', 'monocyte', 'neutrophil', 'platelet', 'RBC', 'WBC', 'other'] cmap_cls = plt.cm.tab10(np.linspace(0, 1, 11)) # Convert PIL image to BGR for OpenCV contour segmentation img_bgr = cv2.cvtColor(raw_image_np, cv2.COLOR_RGB2BGR) h, w = img_bgr.shape[:2] # Run cell detection contours # We use an adaptive threshold to handle small image dimensions (crops) as well area_thresh = 100 if max(h, w) < 400 else 500 boxes = segment_cells(img_bgr, min_area=area_thresh)[:20] # Limit to top 20 cells results = [] for (x, y, wb, hb) in boxes: pad = 5 x1, y1 = max(0, x - pad), max(0, y - pad) x2, y2 = min(img_bgr.shape[1], x + wb + pad), min(img_bgr.shape[0], y + hb + pad) crop = img_bgr[y1:y2, x1:x2] if crop.size == 0: continue # Preprocess crop (128x128, normalized RGB) crop_rgb = cv2.cvtColor(crop, cv2.COLOR_BGR2RGB) crop_128 = cv2.resize(crop_rgb, (128, 128)) crop_np = crop_128.astype(np.float32) / 255.0 input_tensor = np.expand_dims(crop_np, axis=0) # Inference preds = model(input_tensor, training=False) probs = preds[0].numpy() pred_idx = int(np.argmax(probs)) results.append({ 'box': (x1, y1, x2, y2), 'label': WBC_CLASSES_TF[pred_idx], 'conf': float(probs[pred_idx]), 'probs': probs, 'crop': crop }) if results: # Draw bounding boxes and text ann = img_bgr.copy() for res in results: x1, y1, x2, y2 = res['box'] i = WBC_CLASSES_TF.index(res['label']) col = tuple(int(c*255) for c in cmap_cls[i][2::-1]) cv2.rectangle(ann, (x1, y1), (x2, y2), col, 2) txt = f"{res['label']} {res['conf']*100:.1f}%" (tw, th), _ = cv2.getTextSize(txt, cv2.FONT_HERSHEY_SIMPLEX, 0.4, 1) cv2.rectangle(ann, (x1, y1 - th - 6), (x1 + tw + 4, y1), col, -1) cv2.putText(ann, txt, (x1 + 2, y1 - 4), cv2.FONT_HERSHEY_SIMPLEX, 0.4, (255, 255, 255), 1, cv2.LINE_AA) # Convert images to RGB for matplotlib ann_rgb = cv2.cvtColor(ann, cv2.COLOR_BGR2RGB) img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB) # Construct Colab-style matplotlib output figure n = len(results) n_cols = min(n, 5) n_rows = (n + n_cols - 1) // n_cols fig = plt.figure(figsize=(18, 5 + n_rows*3 + 4)) from matplotlib.gridspec import GridSpec gs = GridSpec(3, 2, figure=fig, height_ratios=[5, max(1, n_rows*3), 4], hspace=0.5, wspace=0.3) # Original Smear Plot ax0 = fig.add_subplot(gs[0, 0]) ax0.imshow(img_rgb) ax0.axis('off') ax0.set_title('Original', fontweight='bold') # Bounding Boxes Plot ax1 = fig.add_subplot(gs[0, 1]) ax1.imshow(ann_rgb) ax1.axis('off') ax1.set_title(f'Detected: {n} cells', fontweight='bold') # Cropped Cells Grid sub = gs[1, :].subgridspec(n_rows, n_cols, hspace=0.7, wspace=0.35) for idx_c, res in enumerate(results): r, c = divmod(idx_c, n_cols) ax = fig.add_subplot(sub[r, c]) cr = cv2.cvtColor(cv2.resize(res['crop'], (128, 128)), cv2.COLOR_BGR2RGB) ax.imshow(cr) col_idx = WBC_CLASSES_TF.index(res['label']) col_plt = cmap_cls[col_idx] ax.set_title(f"#{idx_c+1} {res['label']}\n{res['conf']*100:.1f}%", fontsize=8, fontweight='bold', color=col_plt) for sp in ax.spines.values(): sp.set_edgecolor(col_plt) sp.set_linewidth(2) ax.set_xticks([]) ax.set_yticks([]) # Confidence Score Bar Chart ax3 = fig.add_subplot(gs[2, :]) lbls = [r['label'] for r in results] confs = [r['conf']*100 for r in results] bcols = [cmap_cls[WBC_CLASSES_TF.index(l)] for l in lbls] bars = ax3.bar(range(n), confs, color=bcols, edgecolor='black', linewidth=0.5) ax3.set_xticks(range(n)) ax3.set_xticklabels([f"#{i+1}\n{lbls[i]}" for i in range(n)], fontsize=8, rotation=30, ha='right') ax3.set_ylabel('Confidence (%)') ax3.set_ylim(0, 113) ax3.axhline(40.0, color='red', ls='--', lw=1.2, label='Threshold (40%)') ax3.legend(fontsize=9) ax3.grid(axis='y', alpha=0.3) ax3.set_title('Confidence per detected cell', fontweight='bold') for bar, v in zip(bars, confs): ax3.text(bar.get_x() + bar.get_width()/2, v + 1.5, f'{v:.1f}%', ha='center', fontsize=7, fontweight='bold') # Summary Box at bottom counts_cnt = Counter(lbls) summary_str = " ".join([f"{cls}: {count}" for cls, count in sorted(counts_cnt.items())]) fig.text(0.5, 0.005, f"Cell count summary: {summary_str}", ha='center', fontsize=11, fontweight='bold', bbox=dict(boxstyle='round,pad=0.4', facecolor='lightyellow', edgecolor='orange')) plt.suptitle(f'Results — {filename}', fontsize=13, fontweight='bold', y=1.01) # Save figure to in-memory buffer buf = io.BytesIO() plt.savefig(buf, format='jpeg', dpi=120, bbox_inches='tight', facecolor='white') buf.seek(0) # Create base64 representation of the combined plot base64_str = base64.b64encode(buf.read()).decode("utf-8") plt.close(fig) # Set prediction outputs summary_label = ", ".join([f"{count} {cls}" for cls, count in counts_cnt.items()]) predicted_label = f"Detected {len(results)} cells: {summary_label}" confidence = float(np.mean([res['conf'] for res in results])) class_probabilities = {WBC_CLASSES_TF[i]: 0.0 for i in range(len(WBC_CLASSES_TF))} for res in results: class_probabilities[res['label']] += 1.0 for cls in class_probabilities: class_probabilities[cls] /= len(results) # Return the base64 plot directly back to the site return predicted_label, confidence, class_probabilities, base64_str, time.time() - start_time else: # Fallback to single-cell prediction if no cells segmented is_whole_smear = False if not is_whole_smear: img_128 = image.resize((128, 128)) img_np = np.array(img_128, dtype=np.float32) / 255.0 input_tensor = np.expand_dims(img_np, axis=0) preds = model(input_tensor, training=False) probs = preds[0].numpy() pred_idx = int(np.argmax(probs)) confidence = float(probs[pred_idx]) predicted_label = WBC_CLASSES_TF[pred_idx] class_probabilities = {WBC_CLASSES_TF[i]: float(probs[i]) for i in range(len(WBC_CLASSES_TF))} heatmap_img = generate_wbc_gradcam_tf(model, input_tensor, pred_idx, raw_image_np) else: # Preprocess for PyTorch model import matplotlib matplotlib.use('Agg') import matplotlib.pyplot as plt from collections import Counter # Setup class names list in the exact order as PyTorch training CLASS_NAMES = ["Basophil", "Eosinophil", "Erythroblast", "IG", "Lymphocyte", "Monocyte", "Neutrophil", "Platelet"] cmap_cls = plt.cm.tab10(np.linspace(0, 1, 8)) # Convert PIL image to BGR for OpenCV contour segmentation img_bgr = cv2.cvtColor(raw_image_np, cv2.COLOR_RGB2BGR) h, w = img_bgr.shape[:2] # Run cell detection contours # Using min_area=400 since that's what was used in the original notebook area_thresh = 150 if max(h, w) < 400 else 400 boxes = segment_cells(img_bgr, min_area=area_thresh)[:20] # Limit to top 20 cells results = [] transform = get_wbc_transforms() for (x, y, wb, hb) in boxes: pad = 5 x1, y1 = max(0, x - pad), max(0, y - pad) x2, y2 = min(img_bgr.shape[1], x + wb + pad), min(img_bgr.shape[0], y + hb + pad) crop = img_bgr[y1:y2, x1:x2] if crop.size == 0: continue # Preprocess crop (transforms handles resizing to 224x224 and normalization) crop_rgb = cv2.cvtColor(crop, cv2.COLOR_BGR2RGB) crop_pil = Image.fromarray(crop_rgb) input_tensor = transform(crop_pil).unsqueeze(0) # PyTorch Inference with torch.no_grad(): outputs = model(input_tensor) probs = torch.softmax(outputs, dim=1)[0].cpu().numpy() pred_idx = int(np.argmax(probs)) results.append({ 'box': (x1, y1, x2, y2), 'label': CLASS_NAMES[pred_idx], 'conf': float(probs[pred_idx]), 'probs': probs, 'crop': crop }) if results: # Draw bounding boxes and text ann = img_bgr.copy() for res in results: x1, y1, x2, y2 = res['box'] i = CLASS_NAMES.index(res['label']) col = tuple(int(c*255) for c in cmap_cls[i][2::-1]) cv2.rectangle(ann, (x1, y1), (x2, y2), col, 2) txt = f"{res['label']} {res['conf']*100:.1f}%" (tw, th), _ = cv2.getTextSize(txt, cv2.FONT_HERSHEY_SIMPLEX, 0.4, 1) cv2.rectangle(ann, (x1, y1 - th - 6), (x1 + tw + 4, y1), col, -1) cv2.putText(ann, txt, (x1 + 2, y1 - 4), cv2.FONT_HERSHEY_SIMPLEX, 0.4, (255, 255, 255), 1, cv2.LINE_AA) # Convert images to RGB for matplotlib ann_rgb = cv2.cvtColor(ann, cv2.COLOR_BGR2RGB) img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB) # Construct Colab-style matplotlib output figure n = len(results) n_cols = min(n, 5) n_rows = (n + n_cols - 1) // n_cols fig = plt.figure(figsize=(18, 5 + n_rows*3 + 4)) from matplotlib.gridspec import GridSpec gs = GridSpec(3, 2, figure=fig, height_ratios=[5, max(1, n_rows*3), 4], hspace=0.5, wspace=0.3) # Original Smear Plot ax0 = fig.add_subplot(gs[0, 0]) ax0.imshow(img_rgb) ax0.axis('off') ax0.set_title('Original', fontweight='bold') # Bounding Boxes Plot ax1 = fig.add_subplot(gs[0, 1]) ax1.imshow(ann_rgb) ax1.axis('off') ax1.set_title(f'Detected: {n} cells', fontweight='bold') # Cropped Cells Grid sub = gs[1, :].subgridspec(n_rows, n_cols, hspace=0.7, wspace=0.35) for idx_c, res in enumerate(results): r, c = divmod(idx_c, n_cols) ax = fig.add_subplot(sub[r, c]) cr = cv2.cvtColor(cv2.resize(res['crop'], (128, 128)), cv2.COLOR_BGR2RGB) ax.imshow(cr) col_idx = CLASS_NAMES.index(res['label']) col_plt = cmap_cls[col_idx] ax.set_title(f"#{idx_c+1} {res['label']}\n{res['conf']*100:.1f}%", fontsize=8, fontweight='bold', color=col_plt) for sp in ax.spines.values(): sp.set_edgecolor(col_plt) sp.set_linewidth(2) ax.set_xticks([]) ax.set_yticks([]) # Confidence Score Bar Chart ax3 = fig.add_subplot(gs[2, :]) lbls = [r['label'] for r in results] confs = [r['conf']*100 for r in results] bcols = [cmap_cls[CLASS_NAMES.index(l)] for l in lbls] bars = ax3.bar(range(n), confs, color=bcols, edgecolor='black', linewidth=0.5) ax3.set_xticks(range(n)) ax3.set_xticklabels([f"#{i+1}\n{lbls[i]}" for i in range(n)], fontsize=8, rotation=30, ha='right') ax3.set_ylabel('Confidence (%)') ax3.set_ylim(0, 113) ax3.axhline(40.0, color='red', ls='--', lw=1.2, label='Threshold (40%)') ax3.legend(fontsize=9) ax3.grid(axis='y', alpha=0.3) ax3.set_title('Confidence per detected cell', fontweight='bold') for bar, v in zip(bars, confs): ax3.text(bar.get_x() + bar.get_width()/2, v + 1.5, f'{v:.1f}%', ha='center', fontsize=7, fontweight='bold') # Summary Box at bottom counts_cnt = Counter(lbls) summary_str = " ".join([f"{cls}: {count}" for cls, count in sorted(counts_cnt.items())]) fig.text(0.5, 0.005, f"Cell count summary: {summary_str}", ha='center', fontsize=11, fontweight='bold', bbox=dict(boxstyle='round,pad=0.4', facecolor='lightyellow', edgecolor='orange')) plt.suptitle(f'Results — {filename}', fontsize=13, fontweight='bold', y=1.01) # Save figure to in-memory buffer buf = io.BytesIO() plt.savefig(buf, format='jpeg', dpi=120, bbox_inches='tight', facecolor='white') buf.seek(0) # Create base64 representation of the combined plot base64_str = base64.b64encode(buf.read()).decode("utf-8") plt.close(fig) # Set prediction outputs summary_label = ", ".join([f"{count} {cls}" for cls, count in counts_cnt.items()]) predicted_label = f"Detected {len(results)} cells: {summary_label}" confidence = float(np.mean([res['conf'] for res in results])) class_probabilities = {WBC_CLASSES[i]: 0.0 for i in range(len(WBC_CLASSES))} for res in results: class_probabilities[res['label']] += 1.0 for cls in class_probabilities: class_probabilities[cls] /= len(results) # Return the base64 plot directly back to the site return predicted_label, confidence, class_probabilities, base64_str, time.time() - start_time else: # Fallback to single-cell prediction if no cells segmented is_whole_smear = False if not is_whole_smear: transform = get_wbc_transforms() input_tensor = transform(image).unsqueeze(0) with torch.no_grad(): outputs = model(input_tensor) probs = torch.softmax(outputs, dim=1)[0].cpu().numpy() pred_idx = int(np.argmax(probs)) confidence = float(probs[pred_idx]) predicted_label = WBC_CLASSES[pred_idx] class_probabilities = {WBC_CLASSES[i]: float(probs[i]) for i in range(len(WBC_CLASSES))} heatmap_img = generate_wbc_gradcam(model, input_tensor, pred_idx, raw_image_np) elif module_type == "skin_lesion": # Load Skin ViT model processor, model = load_skin_model()