Federated Learning for Healthcare
Training Without Seeing Patient Data
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.
Current Round
0 / 10
Global Model Accuracy
73.1%
Privacy Status
- Hospital 1 Local
- Hospital 2 Local
- Hospital 3 Local
- Global Model
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)
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.
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
- Understand the privacy problem: Why centralized training violates HIPAA/GDPR
- Explain federated learning: How models can be trained without sharing data
- Implement FedAvg: Write the core aggregation algorithm from scratch
- Identify challenges: Non-IID data, communication costs, Byzantine attacks
- Deploy responsibly: Recognize when FL is needed vs when centralized is acceptable