Federated learning: entrenar modelos sin compartir datos
Conceptos de federated learning, patrones de implementación y retos prácticos de entrenar modelos en dispositivos distribuidos preservando la privacidad.

El machine learning tradicional requiere reunir todos los datos de entrenamiento en un solo lugar. El federated learning invierte ese modelo: en lugar de llevar los datos al modelo, llevas el modelo a los datos. Cada dispositivo entrena localmente y solo las actualizaciones del modelo viajan al servidor.
Este enfoque resuelve problemas reales. Las organizaciones de salud no pueden compartir historiales de pacientes. Los bancos no pueden agrupar datos de transacciones. Los usuarios móviles no quieren subir sus patrones de escritura. El federated learning les permite a todos beneficiarse de la inteligencia colectiva sin comprometer la privacidad.
Cómo funciona el federated learning
El bucle central es engañosamente simple: distribuir un modelo, entrenar localmente, agregar actualizaciones, repetir. La complejidad se esconde en los detalles de cada paso.
# ❌ Traditional centralized training
def train_centralized(all_data, model):
# All data must exist in one location
for batch in all_data:
loss = model.forward(batch)
loss.backward()
model.update()
return model
# Problem: all_data contains sensitive records from every user# ✅ Federated learning loop
from typing import List, Dict
import numpy as np
class FederatedServer:
def __init__(self, initial_model: Dict[str, np.ndarray]):
self.global_model = initial_model
self.round_number = 0
def select_clients(
self, available_clients: List[str], fraction: float = 0.1
) -> List[str]:
"""Select a random subset of clients for this round."""
num_selected = max(1, int(len(available_clients) * fraction))
indices = np.random.choice(
len(available_clients), num_selected, replace=False
)
return [available_clients[i] for i in indices]
def aggregate_updates(
self, client_updates: List[Dict[str, np.ndarray]],
client_sizes: List[int]
) -> Dict[str, np.ndarray]:
"""Weighted average of client model updates (FedAvg)."""
total_size = sum(client_sizes)
aggregated = {}
for key in self.global_model:
weighted_sum = sum(
update[key] * (size / total_size)
for update, size in zip(client_updates, client_sizes)
)
aggregated[key] = weighted_sum
self.global_model = aggregated
self.round_number += 1
return aggregatedCada cliente entrena con sus propios datos, calcula la diferencia entre el modelo actualizado y el modelo recibido, y envía solo ese delta de vuelta. El servidor nunca ve datos crudos, solo gradientes o pesos del modelo promediados.
Implementación del entrenamiento en el cliente
El componente cliente maneja el entrenamiento local y comunica solo actualizaciones del modelo. Los datos de entrenamiento originales nunca abandonan el dispositivo.
class FederatedClient:
def __init__(self, client_id: str, local_data: np.ndarray):
self.client_id = client_id
self.local_data = local_data
self.local_model = None
def receive_model(
self, global_model: Dict[str, np.ndarray]
) -> None:
"""Receive the latest global model from server."""
self.local_model = {
k: v.copy() for k, v in global_model.items()
}
def train_local(
self, epochs: int = 5, learning_rate: float = 0.01
) -> Dict[str, np.ndarray]:
"""Train on local data and return model update."""
for epoch in range(epochs):
for batch in self._get_batches(self.local_data):
gradients = self._compute_gradients(batch)
for key in self.local_model:
self.local_model[key] -= learning_rate * gradients[key]
return self.local_model
def _get_batches(
self, data: np.ndarray, batch_size: int = 32
):
"""Split local data into training batches."""
indices = np.random.permutation(len(data))
for i in range(0, len(data), batch_size):
batch_idx = indices[i:i + batch_size]
yield data[batch_idx]
def _compute_gradients(
self, batch: np.ndarray
) -> Dict[str, np.ndarray]:
"""Compute gradients for a single batch."""
# Simplified - real implementation uses autograd
gradients = {}
for key, weights in self.local_model.items():
gradients[key] = np.random.randn(*weights.shape) * 0.01
return gradientsLa cantidad de épocas locales importa significativamente. Con muy pocas, los clientes apenas aprenden de sus datos. Con demasiadas, los modelos de los clientes divergen entre sí, haciendo que la agregación sea menos efectiva: un fenómeno llamado desviación del cliente (client drift).
Manejo de la distribución no IID de los datos
En la práctica, los datos entre clientes raramente están distribuidos de forma idéntica. Un modelo de predicción de teclado en el teléfono de un desarrollador ve patrones distintos que en el de un adolescente. Esta naturaleza no IID (no independiente e idénticamente distribuida) es el mayor desafío del federated learning.
# ❌ Assuming uniform data distribution
def naive_aggregate(updates: List[Dict[str, np.ndarray]]):
"""Simple average assumes all clients have similar data."""
aggregated = {}
for key in updates[0]:
aggregated[key] = np.mean(
[u[key] for u in updates], axis=0
)
return aggregated# ✅ FedProx: Adding a proximal term to handle heterogeneity
class FedProxClient(FederatedClient):
def __init__(
self, client_id: str, local_data: np.ndarray, mu: float = 0.01
):
super().__init__(client_id, local_data)
self.mu = mu # Proximal term strength
self.reference_model = None
def receive_model(
self, global_model: Dict[str, np.ndarray]
) -> None:
super().receive_model(global_model)
self.reference_model = {
k: v.copy() for k, v in global_model.items()
}
def train_local(
self, epochs: int = 5, learning_rate: float = 0.01
) -> Dict[str, np.ndarray]:
"""Train with proximal term to limit client drift."""
for epoch in range(epochs):
for batch in self._get_batches(self.local_data):
gradients = self._compute_gradients(batch)
for key in self.local_model:
# Standard gradient update
update = learning_rate * gradients[key]
# Proximal term: penalize divergence from global model
proximal_penalty = self.mu * (
self.local_model[key] - self.reference_model[key]
)
self.local_model[key] -= update + proximal_penalty
return self.local_modelFedProx agrega un término de regularización que penaliza a los modelos locales por alejarse demasiado del modelo global. El hiperparámetro mu controla el equilibrio: valores más altos mantienen a los clientes más cerca del modelo global, pero pueden subajustar los datos locales.
Garantías de privacidad con differential privacy
El federated learning por sí solo no garantiza privacidad. Las actualizaciones del modelo pueden filtrar información sobre los datos de entrenamiento mediante ataques de inversión de gradientes. Agregar differential privacy proporciona garantías matemáticas.
class DifferentiallyPrivateClient(FederatedClient):
def __init__(
self,
client_id: str,
local_data: np.ndarray,
noise_multiplier: float = 1.0,
max_grad_norm: float = 1.0,
):
super().__init__(client_id, local_data)
self.noise_multiplier = noise_multiplier
self.max_grad_norm = max_grad_norm
def _clip_gradients(
self, gradients: Dict[str, np.ndarray]
) -> Dict[str, np.ndarray]:
"""Clip gradient norm to bound sensitivity."""
total_norm = np.sqrt(
sum(np.sum(g ** 2) for g in gradients.values())
)
clip_factor = min(1.0, self.max_grad_norm / (total_norm + 1e-6))
return {
key: grad * clip_factor
for key, grad in gradients.items()
}
def _add_noise(
self, gradients: Dict[str, np.ndarray]
) -> Dict[str, np.ndarray]:
"""Add calibrated Gaussian noise for differential privacy."""
noise_scale = self.noise_multiplier * self.max_grad_norm
return {
key: grad + np.random.normal(0, noise_scale, grad.shape)
for key, grad in gradients.items()
}
def train_local(
self, epochs: int = 5, learning_rate: float = 0.01
) -> Dict[str, np.ndarray]:
"""Train with DP guarantees: clip then noise."""
for epoch in range(epochs):
for batch in self._get_batches(self.local_data):
gradients = self._compute_gradients(batch)
gradients = self._clip_gradients(gradients)
gradients = self._add_noise(gradients)
for key in self.local_model:
self.local_model[key] -= learning_rate * gradients[key]
return self.local_modelEl equilibrio entre privacidad y utilidad es real. Más ruido significa garantías de privacidad más fuertes pero una convergencia más lenta. En la práctica, debes rastrear el presupuesto de privacidad (épsilon) a lo largo de las rondas y detener el entrenamiento cuando se agote.
Estrategias de eficiencia en la comunicación
El ancho de banda es el cuello de botella en el federated learning. Enviar actualizaciones completas del modelo desde miles de dispositivos es costoso. Las técnicas de compresión reducen los costos de comunicación drásticamente.
def compress_update_top_k(
update: Dict[str, np.ndarray], k_fraction: float = 0.1
) -> Dict[str, tuple]:
"""Keep only top-k% of gradient values by magnitude."""
compressed = {}
for key, values in update.items():
flat = values.flatten()
k = max(1, int(len(flat) * k_fraction))
top_indices = np.argpartition(np.abs(flat), -k)[-k:]
top_values = flat[top_indices]
compressed[key] = (
top_indices,
top_values,
values.shape,
)
return compressed
def decompress_update(
compressed: Dict[str, tuple]
) -> Dict[str, np.ndarray]:
"""Reconstruct full update from compressed representation."""
decompressed = {}
for key, (indices, values, shape) in compressed.items():
full = np.zeros(np.prod(shape))
full[indices] = values
decompressed[key] = full.reshape(shape)
return decompressedLa esparsificación top-k típicamente retiene más del 90% de la calidad del modelo transmitiendo solo del 1% al 10% de los parámetros. Combinada con cuantización (reducir float32 a int8), puedes lograr una compresión de 100x con una pérdida mínima de precisión.
Orquestación del ciclo completo de entrenamiento
Juntar todo requiere una orquestación cuidadosa de la selección de clientes, el entrenamiento, la agregación y la evaluación a lo largo de varias rondas.
def run_federated_training(
server: FederatedServer,
clients: List[FederatedClient],
num_rounds: int = 100,
clients_per_round: float = 0.1,
local_epochs: int = 5,
) -> List[float]:
"""Run complete federated training loop."""
accuracies = []
for round_num in range(num_rounds):
# Step 1: Select participating clients
client_ids = [c.client_id for c in clients]
selected_ids = server.select_clients(client_ids, clients_per_round)
selected_clients = [
c for c in clients if c.client_id in selected_ids
]
# Step 2: Distribute global model
for client in selected_clients:
client.receive_model(server.global_model)
# Step 3: Local training
updates = []
sizes = []
for client in selected_clients:
update = client.train_local(epochs=local_epochs)
updates.append(update)
sizes.append(len(client.local_data))
# Step 4: Aggregate updates
server.aggregate_updates(updates, sizes)
# Step 5: Evaluate
accuracy = evaluate_model(server.global_model)
accuracies.append(accuracy)
if round_num % 10 == 0:
print(
f"Round {round_num}: accuracy={accuracy:.4f}, "
f"clients={len(selected_clients)}"
)
return accuracies
def evaluate_model(
model: Dict[str, np.ndarray]
) -> float:
"""Evaluate model on held-out test set."""
# Simplified evaluation
return 0.0Los sistemas federados en producción agregan tolerancia a fallos (clientes que se desconectan a mitad de una ronda), agregación segura (el servidor no puede ver actualizaciones individuales) y actualizaciones asíncronas (no esperar a los clientes lentos). Cada capa añade complejidad pero resuelve desafíos reales de despliegue.
Conclusiones clave
El federated learning representa un cambio fundamental en cómo pensamos sobre los datos y el entrenamiento de modelos. La idea central es que los modelos pueden aprender de datos que nunca ven directamente, y esto habilita casos de uso que antes eran imposibles debido a regulaciones de privacidad, preocupaciones competitivas o simples limitaciones logísticas.
Los desafíos técnicos son reales: las distribuciones no IID de los datos degradan la calidad del modelo, los costos de comunicación escalan con el tamaño del modelo y la cantidad de clientes, y las garantías de privacidad tienen un costo en la precisión del modelo. Pero el campo madura rápidamente. Si tu aplicación involucra datos sensibles de usuarios y has estado limitado por el paradigma de entrenamiento centralizado, el federated learning podría ser el enfoque que desbloquee tu próximo avance.


