Federated Learning for Healthcare

Training Without Seeing Patient Data

Privacy-Preserving ML
HIPAA Compliant
Distributed Training
The Challenge

Scenario: You want to build a deep learning model to detect heart arrhythmias from ECG data. Hospital A has 10,000 patients, Hospital B has 8,000, Hospital C has 12,000. But they can't share patient data (HIPAA violations, privacy laws, competitive concerns).

Solution: Federated Learning — Train a global model across all hospitals without the data ever leaving each institution.

Federated Training Simulation
Watch 3 hospitals train a shared model collaboratively

Current Round

0 / 10

Global Model Accuracy

73.1%

Privacy Status

High Privacy
Hospital 1
Total Samples:1000
Normal Cases:700
Abnormal Cases:300
Data Distribution:70% / 30%
Local Model Acc:
Not trained
Hospital 2
Total Samples:1000
Normal Cases:500
Abnormal Cases:500
Data Distribution:50% / 50%
Local Model Acc:
Not trained
Hospital 3
Total Samples:1000
Normal Cases:600
Abnormal Cases:400
Data Distribution:60% / 40%
Local Model Acc:
Not trained
Global Model Accuracy Over Training Rounds
How the federated model improves as hospitals share model updates (not data)
0Training Round0.60.70.80.91Accuracy
  • Hospital 1 Local
  • Hospital 2 Local
  • Hospital 3 Local
  • Global Model
Implementation: Federated Averaging Algorithm
Core code for federated learning in PyTorch
import torch
import torch.nn as nn
from typing import List

class FederatedServer:
    def __init__(self, model: nn.Module):
        self.global_model = model
        
    def federated_averaging(self, local_models: List[nn.Module], 
                           local_weights: List[int]):
        """
        FedAvg: Aggregate local models into global model
        weighted by number of samples at each hospital
        """
        total_samples = sum(local_weights)
        
        # Initialize global parameters
        global_dict = self.global_model.state_dict()
        
        # Weighted average of all local model parameters
        for key in global_dict.keys():
            global_dict[key] = torch.zeros_like(global_dict[key])
            
            for local_model, num_samples in zip(local_models, local_weights):
                local_dict = local_model.state_dict()
                # Weight by hospital's sample size
                weight = num_samples / total_samples
                global_dict[key] += local_dict[key] * weight
        
        self.global_model.load_state_dict(global_dict)
        return self.global_model

class Hospital:
    def __init__(self, hospital_id: int, data_loader, model: nn.Module):
        self.id = hospital_id
        self.data_loader = data_loader
        self.local_model = model
        
    def local_training(self, epochs: int = 5):
        """
        Train model locally on hospital's private data
        Data NEVER leaves this hospital's servers
        """
        self.local_model.train()
        optimizer = torch.optim.Adam(self.local_model.parameters())
        criterion = nn.CrossEntropyLoss()
        
        for epoch in range(epochs):
            for batch_x, batch_y in self.data_loader:
                optimizer.zero_grad()
                outputs = self.local_model(batch_x)
                loss = criterion(outputs, batch_y)
                loss.backward()
                optimizer.step()
        
        # Return model updates (not data!)
        return self.local_model

# Federated Learning Training Loop
def federated_learning(hospitals: List[Hospital], 
                      global_model: nn.Module, 
                      rounds: int = 10):
    
    server = FederatedServer(global_model)
    
    for round_num in range(rounds):
        print(f"\nRound {round_num + 1}/{rounds}")
        
        # 1. Broadcast global model to all hospitals
        for hospital in hospitals:
            hospital.local_model.load_state_dict(
                server.global_model.state_dict()
            )
        
        # 2. Each hospital trains locally (in parallel!)
        local_models = []
        local_weights = []
        
        for hospital in hospitals:
            local_model = hospital.local_training(epochs=5)
            local_models.append(local_model)
            local_weights.append(len(hospital.data_loader.dataset))
        
        # 3. Server aggregates updates (Federated Averaging)
        server.federated_averaging(local_models, local_weights)
        
    return server.global_model

# Example: 3 hospitals with ECG data
hospitals = [
    Hospital(1, hospital1_loader, create_ecg_model()),
    Hospital(2, hospital2_loader, create_ecg_model()),
    Hospital(3, hospital3_loader, create_ecg_model())
]

global_model = federated_learning(hospitals, create_ecg_model(), rounds=10)
Graduate-Level Concepts

Non-IID Data Challenge

Each hospital's data distribution is different (different patient demographics, disease prevalence). This makes federated learning harder than centralized training because model updates can conflict.

Federated Averaging (FedAvg)

The core algorithm: hospitals train locally, send only model weights (not data) to server, server computes weighted average based on dataset sizes. Simple but effective.

Differential Privacy

Add noise to model updates before sharing to prevent model inversion attacks. Trade-off: more privacy = slightly lower accuracy. ε-differential privacy quantifies this.

Byzantine-Robust Aggregation

What if a malicious hospital sends poisoned updates? Use robust aggregation (median, trimmed mean) instead of simple averaging to defend against adversarial participants.

Production Deployment Considerations

1. Infrastructure Requirements

  • Secure aggregation server (often cloud-based, e.g., Google Cloud Healthcare API)
  • TLS encryption for model weight transmission
  • Each hospital runs local training on their own GPUs/servers

2. Regulatory Compliance

  • HIPAA: Patient data never leaves hospital → compliant by design
  • GDPR: No personal data transfer → privacy preserved
  • FDA approval: Required for clinical decision support models

3. Communication Efficiency

  • Model updates can be large (e.g., ResNet50 = 98MB per round)
  • Use gradient compression (top-k sparsification, quantization) to reduce bandwidth
  • Federated learning over slow networks = major challenge

4. Real-World Frameworks

TensorFlow Federated (TFF)
PySyft
NVIDIA FLARE
Flower (flwr)
10-Minute Demo Learning Objectives
  1. Understand the privacy problem: Why centralized training violates HIPAA/GDPR
  2. Explain federated learning: How models can be trained without sharing data
  3. Implement FedAvg: Write the core aggregation algorithm from scratch
  4. Identify challenges: Non-IID data, communication costs, Byzantine attacks
  5. Deploy responsibly: Recognize when FL is needed vs when centralized is acceptable

© 2026 Dr. Priyamvada Tripathi. All rights reserved.

You are free to share and adapt this content with attribution for non-commercial purposes under the same license.