Chuyển đến nội dung chính

Lesson 15: Capstone — Building an End-to-End Medical AI Pipeline

Final project: X-ray classification system or Clinical NLP pipeline. From data processing to deployment in compliance with regulations.

🧠 AI & ML — Lesson 14 Lesson 15: Capstone — Building Medical AI Pipeline End-to-End

AI in Health & Healthcare: Real Battle Applications

Part 4: Production & Compliance

xdev.asia

This is the final lesson of the series. You will build a complete Medical AI system from raw DICOM → clinical API — integrating all knowledge from lessons 1 to 14.


Capstone Project: Chest X-ray AI System

Goal: Build end-to-end chest X-ray analysis system:

  1. Ingest DICOM from PACS (simulated)
  2. Preprocess (windowing, normalization)
  3. Classify 14 pathologies (CheXpert)
  4. Generate Grad-CAM heatmaps
  5. Serve via HIPAA-compliant REST API
  6. Monitor with drift detection

Stack: PyTorch + FastAPI + Docker + MLflow


Phase 1: Data Pipeline

import pydicom
import numpy as np
import cv2
from pathlib import Path
import torch
from torchvision import transforms

class ChestXrayPipeline:
    """End-to-end pipeline từ DICOM → tensor cho inference."""
    
    # 14 CheXpert pathologies
    PATHOLOGIES = [
        "No Finding", "Enlarged Cardiomediastinum", "Cardiomegaly",
        "Lung Opacity", "Lung Lesion", "Edema", "Consolidation",
        "Pneumonia", "Atelectasis", "Pneumothorax", "Pleural Effusion",
        "Pleural Other", "Fracture", "Support Devices"
    ]
    
    def __init__(self):
        self.transform = transforms.Compose([
            transforms.ToPILImage(),
            transforms.Resize((320, 320)),
            transforms.ToTensor(),
            transforms.Normalize(mean=[0.485, 0.456, 0.406],
                               std=[0.229, 0.224, 0.225])
        ])
    
    def load_dicom(self, dicom_path: str) -> np.ndarray:
        """Load DICOM và convert to normalized 8-bit grayscale."""
        ds = pydicom.dcmread(dicom_path)
        pixel_array = ds.pixel_array.astype(np.float32)
        
        # Apply rescale slope/intercept
        slope = float(getattr(ds, 'RescaleSlope', 1.0))
        intercept = float(getattr(ds, 'RescaleIntercept', 0.0))
        pixel_array = pixel_array * slope + intercept
        
        # Invert if MONOCHROME1
        photometric = getattr(ds, 'PhotometricInterpretation', 'MONOCHROME2')
        if photometric == 'MONOCHROME1':
            pixel_array = pixel_array.max() - pixel_array
        
        # Normalize to 0-255
        p2, p98 = np.percentile(pixel_array, [2, 98])
        pixel_array = np.clip(pixel_array, p2, p98)
        pixel_array = ((pixel_array - p2) / (p98 - p2) * 255).astype(np.uint8)
        
        # Apply CLAHE
        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))
        pixel_array = clahe.apply(pixel_array)
        
        # Convert to 3-channel
        image_rgb = cv2.cvtColor(pixel_array, cv2.COLOR_GRAY2RGB)
        return image_rgb
    
    def preprocess(self, image: np.ndarray) -> torch.Tensor:
        """Image → normalized tensor."""
        return self.transform(image).unsqueeze(0)

Phase 2: Model Architecture

import torch
import torch.nn as nn
from torchvision.models import densenet121, DenseNet121_Weights

class CheXpertModel(nn.Module):
    """DenseNet-121 fine-tuned cho CheXpert 14 pathologies."""
    
    def __init__(self, n_classes: int = 14, pretrained: bool = True):
        super().__init__()
        
        weights = DenseNet121_Weights.IMAGENET1K_V1 if pretrained else None
        backbone = densenet121(weights=weights)
        
        # Freeze early layers
        layers = list(backbone.features.children())
        for layer in layers[:6]:
            for param in layer.parameters():
                param.requires_grad = False
        
        self.features = backbone.features
        
        in_features = backbone.classifier.in_features
        self.classifier = nn.Sequential(
            nn.AdaptiveAvgPool2d((1, 1)),
            nn.Flatten(),
            nn.Linear(in_features, 512),
            nn.BatchNorm1d(512),
            nn.ReLU(),
            nn.Dropout(0.5),
            nn.Linear(512, n_classes)
        )
        
        # Store last conv layer for Grad-CAM
        self.last_conv = self.features[-1]
        self._activations = None
        self._gradients = None
    
    def forward(self, x: torch.Tensor) -> torch.Tensor:
        features = self.features(x)
        return self.classifier(features)
    
    def register_hooks(self):
        """Register hooks cho Grad-CAM."""
        def save_activation(module, input, output):
            self._activations = output
        
        def save_gradient(module, grad_input, grad_output):
            self._gradients = grad_output[0]
        
        self.last_conv.register_forward_hook(save_activation)
        self.last_conv.register_backward_hook(save_gradient)


def compute_gradcam(model, image_tensor, class_idx):
    """Generate Grad-CAM heatmap."""
    model.register_hooks()
    model.eval()
    
    output = model(image_tensor)
    model.zero_grad()
    output[0, class_idx].backward()
    
    activations = model._activations.detach()
    gradients = model._gradients.detach()
    
    weights = gradients.mean(dim=[2, 3], keepdim=True)
    cam = (weights * activations).sum(dim=1, keepdim=True)
    cam = torch.relu(cam).squeeze()
    
    # Normalize
    cam = cam - cam.min()
    cam = cam / (cam.max() + 1e-8)
    
    # Upsample
    H, W = image_tensor.shape[2:]
    cam_np = cam.cpu().numpy()
    cam_resized = cv2.resize(cam_np, (W, H))
    
    return cam_resized

Phase 3: Training Pipeline

import mlflow
import mlflow.pytorch
from torch.cuda.amp import GradScaler, autocast

def train_chexpert(
    model: CheXpertModel,
    train_loader,
    val_loader,
    n_epochs: int = 30,
    lr: float = 1e-4,
    device: str = "cuda"
):
    """Production training với MLflow tracking."""
    
    with mlflow.start_run(run_name="chexpert-densenet121"):
        # Log hyperparameters
        mlflow.log_params({
            "model": "densenet121",
            "lr": lr,
            "n_epochs": n_epochs,
            "batch_size": train_loader.batch_size,
        })
        
        optimizer = torch.optim.AdamW(
            filter(lambda p: p.requires_grad, model.parameters()),
            lr=lr, weight_decay=1e-5
        )
        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=n_epochs)
        scaler = GradScaler()
        
        # Weighted BCE: upweight positive cases
        pos_weights = compute_positive_weights(train_loader)
        criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weights.to(device))
        
        best_val_auc = 0
        
        for epoch in range(n_epochs):
            model.train()
            train_losses = []
            
            for batch in train_loader:
                images = batch["image"].to(device)
                labels = batch["labels"].float().to(device)
                
                optimizer.zero_grad()
                with autocast():
                    logits = model(images)
                    loss = criterion(logits, labels)
                
                scaler.scale(loss).backward()
                scaler.unscale_(optimizer)
                nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
                scaler.step(optimizer)
                scaler.update()
                train_losses.append(loss.item())
            
            # Validation
            val_metrics = evaluate_chexpert(model, val_loader, device)
            scheduler.step()
            
            mlflow.log_metrics({
                "train_loss": np.mean(train_losses),
                "val_auc_mean": val_metrics["mean_auc"],
                **{f"val_auc_{p}": v for p, v in val_metrics["per_class_auc"].items()}
            }, step=epoch)
            
            # Save best model
            if val_metrics["mean_auc"] > best_val_auc:
                best_val_auc = val_metrics["mean_auc"]
                mlflow.pytorch.log_model(model, "best_model")
                print(f"Epoch {epoch}: New best AUC = {best_val_auc:.4f}")
        
        mlflow.log_metric("best_val_auc", best_val_auc)
        print(f"Training complete. Best val AUC: {best_val_auc:.4f}")


def compute_positive_weights(loader) -> torch.Tensor:
    """Compute class weights for imbalanced multilabel dataset."""
    label_counts = None
    total = 0
    for batch in loader:
        labels = batch["labels"].numpy()
        if label_counts is None:
            label_counts = labels.sum(axis=0)
        else:
            label_counts += labels.sum(axis=0)
        total += labels.shape[0]
    
    pos_weight = (total - label_counts) / (label_counts + 1e-6)
    return torch.tensor(pos_weight, dtype=torch.float32)

Phase 4: Evaluation

from sklearn.metrics import roc_auc_score, average_precision_score
import pandas as pd

def evaluate_chexpert(model, loader, device):
    model.eval()
    all_labels = []
    all_probs = []
    
    with torch.no_grad():
        for batch in loader:
            images = batch["image"].to(device)
            labels = batch["labels"].numpy()
            
            logits = model(images)
            probs = torch.sigmoid(logits).cpu().numpy()
            
            all_labels.append(labels)
            all_probs.append(probs)
    
    all_labels = np.concatenate(all_labels)
    all_probs = np.concatenate(all_probs)
    
    per_class_auc = {}
    for i, pathology in enumerate(ChestXrayPipeline.PATHOLOGIES):
        if all_labels[:, i].sum() > 0:
            auc = roc_auc_score(all_labels[:, i], all_probs[:, i])
            per_class_auc[pathology] = round(float(auc), 4)
    
    # Clinical metrics per pathology at threshold=0.5
    threshold = 0.5
    preds = (all_probs >= threshold).astype(int)
    
    clinical_report = []
    for i, pathology in enumerate(ChestXrayPipeline.PATHOLOGIES):
        if all_labels[:, i].sum() > 0:
            tp = ((preds[:, i] == 1) & (all_labels[:, i] == 1)).sum()
            tn = ((preds[:, i] == 0) & (all_labels[:, i] == 0)).sum()
            fp = ((preds[:, i] == 1) & (all_labels[:, i] == 0)).sum()
            fn = ((preds[:, i] == 0) & (all_labels[:, i] == 1)).sum()
            
            sensitivity = tp / (tp + fn + 1e-8)
            specificity = tn / (tn + fp + 1e-8)
            
            clinical_report.append({
                "pathology": pathology,
                "sensitivity": round(float(sensitivity), 4),
                "specificity": round(float(specificity), 4),
                "auc": per_class_auc.get(pathology, 0)
            })
    
    return {
        "per_class_auc": per_class_auc,
        "mean_auc": round(float(np.mean(list(per_class_auc.values()))), 4),
        "clinical_report": pd.DataFrame(clinical_report)
    }

Phase 5: API Integration (from Lesson 14)

from fastapi import FastAPI
import mlflow.pytorch
import io
from PIL import Image

app = FastAPI(title="CheXpert AI API v1.0")

# Load model từ MLflow
MODEL_URI = "runs:/best_run_id/best_model"
ai_model = mlflow.pytorch.load_model(MODEL_URI)
ai_model.eval()

pipeline = ChestXrayPipeline()

@app.post("/api/v1/analyze-xray")
async def analyze_xray(dicom_path: str, generate_heatmap: bool = True):
    """Analyze chest X-ray DICOM file."""
    # Load & preprocess
    image = pipeline.load_dicom(dicom_path)
    tensor = pipeline.preprocess(image).to("cuda")
    
    # Inference
    with torch.no_grad():
        logits = ai_model(tensor)
        probs = torch.sigmoid(logits).squeeze().cpu().numpy()
    
    # Format findings
    findings = []
    for i, (pathology, prob) in enumerate(zip(ChestXrayPipeline.PATHOLOGIES, probs)):
        if prob > 0.3:  # Threshold: tunable
            finding = {
                "pathology": pathology,
                "confidence": round(float(prob), 4),
                "severity": "high" if prob > 0.7 else "moderate" if prob > 0.5 else "low"
            }
            # Generate Grad-CAM for this finding
            if generate_heatmap:
                finding["heatmap"] = compute_gradcam(ai_model, tensor, i).tolist()
            findings.append(finding)
    
    return {
        "findings": sorted(findings, key=lambda x: x["confidence"], reverse=True),
        "total_pathologies_detected": len(findings),
        "disclaimer": "For clinical decision support only. Not for standalone diagnosis.",
    }

Conclusion Series

Congratulations — you have completed the AI in Health & Healthcare series!

What you learned:

PartTopicsSkills
Part 1PlatformDICOM, FHIR, medical data preprocessing
Part 2Medical ImagingCNN, U-Net, YOLO, WSI analysis
Part 3Clinical AINLP, Drug Discovery GNN, Genomics
Part 4ProductionFederated Learning, XAI, FDA, MLOps
CapstoneEnd-to-endFull system from DICOM to API

Next step:

  1. Kaggle competitions: RSNA Pneumonia, PadChest, CheXpert
  2. Datasets: MIMIC-IV, PhysioNet, UK Biobank
  3. Conferences: MICCAI, NeurIPS Medical Imaging Workshop
  4. Certifications: AWS Healthcare, GCP Healthcare API
  5. Open source: contribute to monai, nnU-Net, medperf

Final Capstone Assignment

Minimum requirements (to be considered completed series):

  1. Train CheXpert model to mean AUC ≥ 0.80 on validation set.
  2. Deploy API with Docker. Test with 10 X-ray DICOM samples.
  3. Generate Grad-CAM for each prediction. Overlay to original image.
  4. Add drift detection: compare prediction confidence distribution after 100 requests vs validation set.
  5. Write a simple model card (1 page): intended use, performance metrics, limitations, ethical considerations.

Bonus (advanced):

  • Fine-tune with ViT-B/16 (Vision Transformer). Compare with DenseNet-121.
  • Implement federated version: simulate 3 hospitals with CheXpert subset.
  • Submit results to CheXpert leaderboard.