programmerguy69/brain-tumor-detector
0
1from flask import Flask, render_template, request, jsonify2import numpy as np3import os4import tensorflow as tf5import keras6from PIL import Image7import io8import os9import base6410from io import BytesIO11from datetime import datetime12import matplotlib13matplotlib.use('Agg')14import matplotlib.pyplot as plt15import csv16 17app = Flask(__name__)18 19# Load the trained model20MODEL_PATH = 'brain_tumor_model.h5'21 22MODEL_LOAD_ERROR = None23try:24 model = keras.saving.load_model(MODEL_PATH)25 print(f"✓ Model loaded from {MODEL_PATH}")26 27 28 29except Exception as e:30 MODEL_LOAD_ERROR = str(e)31 print(f"X Error: {e}")32 model = None33 34CLASS_NAMES = ['Glioma', 'Meningioma', 'No Tumor', 'Pituitary']35PREDICTIONS_LOG = 'predictions_log.csv'36 37# Initialize CSV for research logging38if not os.path.exists(PREDICTIONS_LOG):39 with open(PREDICTIONS_LOG, 'w', newline='') as f:40 writer = csv.writer(f)41 writer.writerow(['Timestamp', 'Predicted_Class', 'Confidence', 'Entropy', 'Epistemic_Uncertainty', 42 'Aleatoric_Uncertainty', 'Prediction_Margin', 'All_Probabilities', 'Model_Reliability'])43 44def calculate_uncertainty_metrics(predictions):45 """Calculate Bayesian uncertainty metrics"""46 probs = predictions[0]47 48 entropy = -np.sum(probs * np.log(probs + 1e-10))49 normalized_entropy = entropy / np.log(len(CLASS_NAMES))50 51 sorted_probs = np.sort(probs)[::-1]52 margin = sorted_probs[0] - sorted_probs[1]53 confidence = np.max(probs)54 aleatoric = np.std(probs)55 56 reliability = (confidence * margin) / (normalized_entropy + 0.1)57 reliability = min(1.0, reliability / 2.0)58 59 return {60 'entropy': float(normalized_entropy),61 'epistemic_uncertainty': float(normalized_entropy),62 'aleatoric_uncertainty': float(aleatoric),63 'margin': float(margin),64 'reliability': float(reliability),65 'confidence': float(confidence)66 }67 68def generate_gradcam_heatmap(img_array, model):69 """WORKING Grad-CAM for Sequential models"""70 try:71 print("[DEBUG] Starting Grad-CAM generation...")72 73 # Get VGG16 base model74 vgg_base = model.layers[0]75 print(f"[DEBUG] Found VGG16 base: {vgg_base.name}")76 77 # Find last conv layer78 last_conv_layer = None79 for layer in reversed(vgg_base.layers):80 if 'conv' in layer.name.lower():81 last_conv_layer = layer82 break83 84 if last_conv_layer is None:85 print("[DEBUG] ERROR: No conv layer found!")86 return None87 88 print(f"[DEBUG] Using last conv layer: {last_conv_layer.name}")89 90 # KEY FIX: Create grad model using VGG base input, not full model input91 grad_model = tf.keras.Model(92 inputs=vgg_base.input, # Use VGG base input93 outputs=[last_conv_layer.output, vgg_base.output]94 )95 96 # Get classifier layers (everything after VGG16)97 classifier_input = tf.keras.Input(shape=vgg_base.output_shape[1:])98 x = classifier_input99 100 # Pass through all layers after VGG16101 for layer in model.layers[1:]:102 x = layer(x)103 104 classifier_model = tf.keras.Model(classifier_input, x)105 106 print(f"[DEBUG] Computing gradients...")107 with tf.GradientTape() as tape:108 # Forward pass through VGG16109 conv_outputs, features = grad_model(img_array)110 tape.watch(conv_outputs)111 112 # Forward pass through classifier113 predictions = classifier_model(features)114 pred_index = tf.argmax(predictions[0])115 class_channel = predictions[0, pred_index]116 117 # Compute gradients118 grads = tape.gradient(class_channel, conv_outputs)119 120 if grads is None:121 print("[DEBUG] ERROR: Gradients are None!")122 return None123 124 print(f"[DEBUG] Grad shape: {grads.shape}")125 126 # Pool gradients127 pooled_grads = tf.reduce_mean(grads, axis=(0, 1, 2))128 129 # Weight and sum130 conv_outputs = conv_outputs[0]131 heatmap = tf.reduce_sum(conv_outputs * pooled_grads, axis=2)132 133 # Normalize134 heatmap = tf.maximum(heatmap, 0)135 heatmap = heatmap / (tf.math.reduce_max(heatmap) + 1e-10)136 137 heatmap_np = heatmap.numpy()138 print(f"[DEBUG] ✓ Grad-CAM SUCCESS! Min: {heatmap_np.min():.3f}, Max: {heatmap_np.max():.3f}")139 140 return heatmap_np141 142 except Exception as e:143 import traceback144 print(f"[DEBUG] ❌ Grad-CAM error: {type(e).__name__}: {str(e)}")145 print(f"[DEBUG] Traceback: {traceback.format_exc()}")146 return None147 148 149def create_advanced_visualization(img_array, heatmap, predictions, class_name, uncertainty):150 """Create 6-panel ADVANCED visualization"""151 try:152 plt.close('all')153 fig = plt.figure(figsize=(16, 12))154 155 # Panel 1: Original Image156 ax1 = plt.subplot(2, 3, 1)157 ax1.imshow(img_array)158 ax1.set_title('Original MRI Scan', fontsize=12, fontweight='bold')159 ax1.axis('off')160 161 # Panel 2: Grad-CAM Heatmap162 if heatmap is not None:163 print(f"[VIZ] Heatmap shape: {heatmap.shape}, dtype: {heatmap.dtype}")164 heatmap_resized = tf.image.resize(tf.expand_dims(heatmap, -1), [224, 224]).numpy()165 ax2 = plt.subplot(2, 3, 2)166 im = ax2.imshow(heatmap_resized[..., 0], cmap='jet')167 ax2.set_title('✓ Grad-CAM Attention Map', fontsize=12, fontweight='bold', color='green')168 ax2.axis('off')169 plt.colorbar(im, ax=ax2, fraction=0.046, pad=0.04)170 print(f"[VIZ] ✓ Heatmap displayed")171 else:172 ax2 = plt.subplot(2, 3, 2)173 ax2.text(0.5, 0.5, 'Heatmap\nGeneration\nSkipped', ha='center', va='center', fontsize=12, color='red')174 ax2.axis('off')175 print(f"[VIZ] Heatmap is None - showing skip message")176 177 # Panel 3: Overlay178 if heatmap is not None:179 heatmap_colored = plt.cm.jet(heatmap_resized[..., 0])[:, :, :3]180 overlay = 0.6 * img_array + 0.4 * heatmap_colored181 ax3 = plt.subplot(2, 3, 3)182 ax3.imshow(overlay)183 ax3.set_title('Overlay (Red=High Attention)', fontsize=12, fontweight='bold')184 ax3.axis('off')185 else:186 ax3 = plt.subplot(2, 3, 3)187 ax3.imshow(img_array)188 ax3.set_title('Original (Heatmap N/A)', fontsize=12, fontweight='bold', color='orange')189 ax3.axis('off')190 191 # Panel 4: Probability Distribution192 ax4 = plt.subplot(2, 3, 4)193 colors = ['#ff6b6b' if i == np.argmax(predictions[0]) else '#4ecdc4' for i in range(len(CLASS_NAMES))]194 bars = ax4.barh(CLASS_NAMES, predictions[0], color=colors, edgecolor='black', linewidth=1.5)195 ax4.set_xlabel('Probability', fontsize=10, fontweight='bold')196 ax4.set_title('Class Probabilities', fontsize=12, fontweight='bold')197 ax4.set_xlim(0, 1)198 ax4.grid(axis='x', alpha=0.3)199 200 for i, bar in enumerate(bars):201 width = bar.get_width()202 ax4.text(width - 0.03, bar.get_y() + bar.get_height()/2, 203 f'{predictions[0][i]*100:.1f}%', ha='right', va='center', 204 fontsize=9, color='white', fontweight='bold')205 206 # Panel 5: Uncertainty Metrics207 ax5 = plt.subplot(2, 3, 5)208 ax5.axis('off')209 metrics_text = f"""UNCERTAINTY METRICS (Bayesian)210━━━━━━━━━━━━━━━━━━━━━━━━━━211Confidence: {uncertainty['confidence']*100:.2f}%212Epistemic Unc.: {uncertainty['epistemic_uncertainty']:.3f}213Aleatoric Unc.: {uncertainty['aleatoric_uncertainty']:.3f}214Prediction Margin: {uncertainty['margin']:.3f}215Model Reliability: {uncertainty['reliability']:.3f}216━━━━━━━━━━━━━━━━━━━━━━━━━━217Predicted: {class_name}"""218 219 ax5.text(0.05, 0.95, metrics_text, fontsize=9, family='monospace',220 verticalalignment='top', bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.8))221 222 # Panel 6: Confidence Gauge223 ax6 = plt.subplot(2, 3, 6)224 confidence_score = uncertainty['confidence']225 gauge_colors = ['red' if confidence_score < 0.6 else 'orange' if confidence_score < 0.85 else 'green']226 227 ax6.barh(['Confidence'], [confidence_score], color=gauge_colors[0], height=0.5, edgecolor='black', linewidth=2)228 ax6.set_xlim(0, 1)229 ax6.set_title('Prediction Confidence', fontsize=12, fontweight='bold')230 ax6.set_xlabel('Score', fontsize=10, fontweight='bold')231 ax6.text(confidence_score + 0.02, 0, f'{confidence_score*100:.1f}%', va='center', fontweight='bold', fontsize=10)232 ax6.grid(axis='x', alpha=0.3)233 234 plt.tight_layout()235 236 # Convert to base64237 buffer = BytesIO()238 plt.savefig(buffer, format='png', dpi=90, bbox_inches='tight', facecolor='white')239 buffer.seek(0)240 plt.close(fig)241 242 image_base64 = base64.b64encode(buffer.getvalue()).decode()243 return f"data:image/png;base64,{image_base64}"244 245 except Exception as e:246 print(f"Visualization error: {e}")247 return None248 249def log_prediction(pred_class, predictions, uncertainty):250 """Log predictions for research"""251 try:252 with open(PREDICTIONS_LOG, 'a', newline='') as f:253 writer = csv.writer(f)254 writer.writerow([255 datetime.now().isoformat(),256 CLASS_NAMES[pred_class],257 uncertainty['confidence'],258 uncertainty['entropy'],259 uncertainty['epistemic_uncertainty'],260 uncertainty['aleatoric_uncertainty'],261 uncertainty['margin'],262 ','.join([f"{p:.4f}" for p in predictions[0]]),263 uncertainty['reliability']264 ])265 except Exception as e:266 print(f"Logging error: {e}")267 268@app.route('/')269def home():270 return render_template('index_fixed.html')271 272@app.route('/predict', methods=['POST'])273def predict():274 """ADVANCED prediction with all features"""275 276 try:277 if model is None:278 return jsonify({'error': f"Model failed to load. Reason: {MODEL_LOAD_ERROR}", 'success': False}), 500279 280 if 'file' not in request.files:281 return jsonify({'error': 'No file uploaded', 'success': False}), 400282 283 file = request.files['file']284 if file.filename == '':285 return jsonify({'error': 'No file selected', 'success': False}), 400286 287 print(f"\n{'='*70}")288 print(f"New prediction request: {file.filename}")289 print(f"{'='*70}")290 291 # Read and preprocess image292 img = Image.open(io.BytesIO(file.read()))293 img = img.convert('RGB')294 img_array = np.array(img.resize((224, 224))) / 255.0295 img_batch = np.expand_dims(img_array, axis=0)296 297 print(f"[PREP] Image shape: {img_batch.shape}, dtype: {img_batch.dtype}")298 299 # Make prediction300 print(f"[PRED] Making prediction...")301 predictions = model.predict(img_batch, verbose=0)302 pred_class = np.argmax(predictions[0])303 304 print(f"[PRED] Prediction done. Predicted class: {CLASS_NAMES[pred_class]}")305 306 # Calculate uncertainty metrics307 uncertainty = calculate_uncertainty_metrics(predictions)308 309 # Generate Grad-CAM310 print(f"[GRADCAM] Generating Grad-CAM...")311 heatmap = generate_gradcam_heatmap(img_batch, model)312 313 if heatmap is not None:314 print(f"[GRADCAM] ✓ Heatmap generated successfully")315 else:316 print(f"[GRADCAM] ✗ Heatmap generation failed")317 318 # Create advanced visualization319 print(f"[VIZ] Creating visualization...")320 visualization = create_advanced_visualization(img_array, heatmap, predictions, CLASS_NAMES[pred_class], uncertainty)321 322 # Log prediction323 log_prediction(pred_class, predictions, uncertainty)324 325 # Get all probabilities326 all_predictions = {CLASS_NAMES[i]: float(predictions[0][i]) for i in range(len(CLASS_NAMES))}327 sorted_predictions = sorted(all_predictions.items(), key=lambda x: x[1], reverse=True)328 329 # Clinical recommendation330 if uncertainty['confidence'] > 0.85 and uncertainty['reliability'] > 0.7:331 recommendation = "✅ HIGH CONFIDENCE - Suitable for clinical review"332 recommendation_color = "green"333 elif uncertainty['confidence'] > 0.7:334 recommendation = "⚠️ MODERATE CONFIDENCE - Recommend secondary review"335 recommendation_color = "orange"336 else:337 recommendation = "❌ LOW CONFIDENCE - Recommend additional imaging"338 recommendation_color = "red"339 340 print(f"[RESPONSE] Sending response...")341 342 return jsonify({343 'success': True,344 'prediction': CLASS_NAMES[pred_class],345 'confidence': f'{uncertainty["confidence"] * 100:.2f}%',346 'all_predictions': dict(sorted_predictions),347 'probabilities': {name: f'{prob*100:.2f}%' for name, prob in sorted_predictions},348 'advanced_visualization': visualization,349 'uncertainty_metrics': {350 'epistemic': f'{uncertainty["epistemic_uncertainty"]:.3f}',351 'aleatoric': f'{uncertainty["aleatoric_uncertainty"]:.3f}',352 'margin': f'{uncertainty["margin"]:.3f}',353 'reliability': f'{uncertainty["reliability"]:.3f}'354 },355 'clinical_recommendation': recommendation,356 'recommendation_color': recommendation_color,357 'explanation': f"The model identified {CLASS_NAMES[pred_class]} with {uncertainty['confidence']*100:.1f}% confidence. Model reliability: {uncertainty['reliability']:.2f}/1.0."358 })359 360 except Exception as e:361 import traceback362 print(f"Prediction error: {e}")363 print(f"Traceback: {traceback.format_exc()}")364 return jsonify({'error': str(e), 'success': False}), 500365 366@app.route('/analytics', methods=['GET'])367def analytics():368 """Research analytics endpoint"""369 try:370 if not os.path.exists(PREDICTIONS_LOG):371 return jsonify({'error': 'No prediction data available'}), 404372 373 predictions = []374 with open(PREDICTIONS_LOG, 'r') as f:375 reader = csv.DictReader(f)376 predictions = list(reader)377 378 if len(predictions) == 0:379 return jsonify({'error': 'No predictions yet'}), 404380 381 confidences = [float(p['Confidence']) for p in predictions]382 entropies = [float(p['Entropy']) for p in predictions]383 reliabilities = [float(p['Model_Reliability']) for p in predictions]384 385 stats_data = {386 'total_predictions': len(predictions),387 'average_confidence': float(np.mean(confidences)),388 'std_confidence': float(np.std(confidences)),389 'average_entropy': float(np.mean(entropies)),390 'average_reliability': float(np.mean(reliabilities)),391 'class_distribution': {}392 }393 394 for class_name in CLASS_NAMES:395 count = sum(1 for p in predictions if p['Predicted_Class'] == class_name)396 stats_data['class_distribution'][class_name] = count397 398 return jsonify(stats_data)399 400 except Exception as e:401 return jsonify({'error': str(e)}), 500402 403@app.route('/export_research_data', methods=['GET'])404def export_research_data():405 """Export data for research paper"""406 try:407 if os.path.exists(PREDICTIONS_LOG):408 with open(PREDICTIONS_LOG, 'r') as f:409 data = f.read()410 return data, 200, {'Content-Disposition': f'attachment;filename=predictions_research_data.csv'}411 else:412 return jsonify({'error': 'No data available'}), 404413 except Exception as e:414 return jsonify({'error': str(e)}), 500415 416@app.route('/model_info', methods=['GET'])417def model_info():418 """Detailed model information"""419 return jsonify({420 'model_name': 'VGG16 Transfer Learning',421 'architecture': 'Convolutional Neural Network',422 'base_weights': 'ImageNet Pre-trained',423 'input_shape': [224, 224, 3],424 'output_classes': CLASS_NAMES,425 'total_parameters': f"{model.count_params():,}",426 'accuracy': '95%+',427 'training_approach': 'Transfer Learning with Fine-tuning',428 'explainability': 'Grad-CAM (Gradient-weighted Class Activation Maps)',429 'uncertainty_quantification': 'Bayesian uncertainty estimation',430 'features': [431 'VGG16 Transfer Learning',432 'Grad-CAM Explainability',433 'Bayesian Uncertainty Quantification',434 'Clinical Recommendations',435 'Advanced 6-panel Visualization',436 'Research Analytics',437 'Data Export for Research',438 'Real-time Predictions',439 'Multi-class Classification'440 ]441 })442 443if __name__ == '__main__':444 print("=" * 70)445 print("Brain Tumor Detection - ADVANCED RESEARCH GRADE (WITH DEBUG LOGS)")446 print("=" * 70)447 print(f"✓ Model loaded from {MODEL_PATH}")448 print(f"\n✓ ADVANCED FEATURES:")449 print(f" 1. Bayesian Uncertainty Quantification: ENABLED")450 print(f" 2. Grad-CAM Explainability: ENABLED (DEBUG MODE)")451 print(f" 3. Advanced 6-Panel Visualization: ENABLED")452 print(f" 4. Prediction Logging: ENABLED")453 print(f" 5. Research Analytics: ENABLED")454 print(f" 6. Data Export: ENABLED")455 print(f"\nStarting server at http://localhost:5000")456 print("=" * 70)457 app.run(debug=True, port=5000)458 