Saltar al contenido

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.

7 min de lectura
Diagrama que muestra federated learning con varios dispositivos aportando actualizaciones de modelo a un servidor central

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.

pypython
# ❌ 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
pypython
# ✅ 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 aggregated

Cada 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.

pypython
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 gradients

La 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.

pypython
# ❌ 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
pypython
# ✅ 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_model

FedProx 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.

pypython
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_model

El 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.

pypython
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 decompressed

La 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.

pypython
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.0

Los 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.

Wilfredo Rujel

Wilfredo Rujel

Ingeniero de Software Full Stack

Compartir esta publicaciónX