ChestSense / CLASS_B_BIAS_MITIGATION.md
NoumanUsman's picture
Upload folder using huggingface_hub
52e8264 verified
|
Raw
History Blame Contribute Delete
4.83 kB

Quick Mitigation Options for Class B Bias

This document outlines immediate workarounds for the Class B bias while you retrain the model.

Option 1: Logit-Bias Correction (Inference-Time Hack)

Add this function to main.py in the /analyze endpoint before softmax:

def apply_logit_bias_correction(logits, bias_correction=None):
    """
    Apply bias correction to logits before softmax.
    This temporarily mitigates the Class B overprediction.
    
    Args:
        logits: Raw model output (before softmax)
        bias_correction: Dict with per-class bias adjustments
    
    Returns:
        Adjusted logits
    """
    if bias_correction is None:
        # DEFAULT: Reduce Class B, boost others
        bias_correction = {
            0: +0.3,   # Adenocarcinoma: slight boost
            1: -1.2,   # Small Cell: heavy penalty
            2: +0.3,   # Large Cell: slight boost
            3: +0.3,   # Squamous Cell: slight boost
        }
    
    for class_idx, bias in bias_correction.items():
        logits[class_idx] += bias
    
    return logits

Usage in classification code:

# in the /analyze endpoint, after model inference:
with torch.no_grad():
    logits = model(input_tensor)  # Raw output
    logits_corrected = apply_logit_bias_correction(logits[0])  # Apply correction
    probs = torch.softmax(logits_corrected, dim=0)  # Then softmax

Option 2: Confidence-Based Filtering

Add confidence thresholds per-class:

CONFIDENCE_THRESHOLDS = {
    0: 0.30,  # Adenocarcinoma: accept if >= 30%
    1: 0.60,  # Small Cell: require higher confidence (70%+)
    2: 0.25,  # Large Cell: accept lower confidence
    3: 0.25,  # Squamous Cell: accept lower confidence
}

if confidence < CONFIDENCE_THRESHOLDS[class_idx]:
    logger.warning(f"LOW CONFIDENCE: {class_idx} at {confidence*100:.1f}% (threshold: {CONFIDENCE_THRESHOLDS[class_idx]*100}%)")
    # Optionally flag as uncertain

Option 3: Report Confidence Scores

Return all confidence scores to the frontend and let the UI handle the ambiguity:

API Response Structure:

{
  "success": true,
  "detections": [
    {
      "tumor_id": 1,
      "prediction": "Small Cell (Class B)",
      "confidence": 75.97,
      "all_confidences": {
        "Adenocarcinoma": 3.5,
        "Small Cell": 75.97,
        "Large Cell": 6.89,
        "Squamous Cell": 13.64
      },
      "confidence_status": "UNCERTAIN",  // NEW: Add this
      "note": "Low confidence in this prediction. Multiple classes viable."
    }
  ]
}

Then in Flutter UI:

// Show warning if confidence is borderline
if (allConfidences.values.any((c) => c > 20 && c != maxConfidence)) {
  showConfidenceWarning("Multiple classes possible. Review with specialist.");
}

Option 4: Ensemble Multiple Models

If you have other trained models:

def ensemble_predict(input_tensor, models_list):
    """Average predictions from multiple models"""
    all_probs = []
    
    for model in models_list:
        with torch.no_grad():
            probs = torch.softmax(model(input_tensor), dim=1)[0]
        all_probs.append(probs)
    
    ensemble_probs = torch.stack(all_probs).mean(dim=0)
    return ensemble_probs

Comparison of Options

Option Ease Effectiveness Drawbacks
Logit Bias Correction ⭐⭐⭐ Easy ⭐⭐ Moderate Doesn't fix root issue; magic numbers
Confidence Filtering ⭐⭐⭐ Easy ⭐⭐ Moderate Still wrong predictions, just flagged
Report All Scores ⭐⭐⭐ Easy ⭐⭐ Moderate Defers to human; adds complexity
Ensemble Models ⭐⭐ Medium ⭐⭐⭐ Good Need multiple models
Retrain with Weights ⭐ Hard ⭐⭐⭐⭐⭐ Excellent Takes time; need training data

⚠️ Recommendation

Use Option 3 (Report All Scores) + Option 1 (Logit Correction) as SHORT-TERM fix

  • Apply light logit correction to reduce Class B dominance
  • Return all confidence scores to frontend
  • Show warnings when confidence is low (< 60%)
  • Plan to retrain model properly (highest priority)

Then retrain the model (Option 5 - see retrain_with_balanced_weights.py)

Testing Your Fix

After implementing mitigation:

# Test with images from each class
LABELS=("Class_A" "Class_B" "Class_E" "Class_E")

for label in "${LABELS[@]}"; do
    img=$(ls test-images/${label}_*.png | head -1)
    echo "Testing: $img"
    curl -X POST -F "file=@$img" http://localhost:5001/analyze | grep -A 10 '"prediction"'
done

Expected results after fix:

Class A image → Mostly Class A (not Class B)
Class B image → Mostly Class B (correct)
Class E image → Mix of E and others (not all B)
Class G image → Mix of G and others (not all B)