Zum Inhalt springen

Federated Learning: Modelle trainieren, ohne Daten zu teilen

Grundlagen des Federated Learning, Implementierungsmuster und praktische Herausforderungen beim Training auf verteilten Geräten unter Datenschutz.

6 Min. Lesezeit
Diagramm, das Federated Learning zeigt: mehrere Geräte liefern Modell-Updates an einen zentralen Server

Traditionelles Machine Learning erfordert, dass alle Trainingsdaten an einem Ort gesammelt werden. Federated Learning kehrt dieses Modell um: Statt die Daten zum Modell zu bringen, bringst du das Modell zu den Daten. Jedes Gerät trainiert lokal, und nur Modell-Updates gelangen zum Server.

Dieser Ansatz löst echte Probleme. Gesundheitsorganisationen können Patientenakten nicht teilen. Banken können Transaktionsdaten nicht zusammenführen. Mobile Nutzer laden ihre Tippmuster nicht hoch. Federated Learning ermöglicht es allen, von kollektiver Intelligenz zu profitieren, ohne die Privatsphäre zu gefährden.

Wie Federated Learning funktioniert

Die Kernschleife ist täuschend einfach: Modell verteilen, lokal trainieren, Updates aggregieren, wiederholen. Die Komplexität steckt im Detail jedes Schritts.

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

Jeder Client trainiert mit seinen eigenen Daten, berechnet die Differenz zwischen dem aktualisierten Modell und dem empfangenen Modell und sendet nur dieses Delta zurück. Der Server sieht niemals Rohdaten — nur gemittelte Gradienten oder Modellgewichte.

Implementierung des clientseitigen Trainings

Die Client-Komponente übernimmt das lokale Training und kommuniziert nur Modell-Updates. Die Roh-Trainingsdaten verlassen das Gerät nie.

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

Die Anzahl lokaler Epochen ist entscheidend. Zu wenige Epochen und die Clients lernen kaum aus ihren Daten. Zu viele Epochen und die Client-Modelle divergieren voneinander, wodurch die Aggregation weniger effektiv wird — ein Phänomen, das man Client-Drift nennt.

Umgang mit nicht-IID Datenverteilung

In der Praxis sind Daten über Clients selten identisch verteilt. Ein Tastaturvorhersagemodell auf dem Telefon eines Entwicklers sieht andere Muster als auf dem eines Teenagers. Diese nicht-IID-Natur (nicht unabhängig und identisch verteilt) ist die größte Herausforderung beim 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 fügt einen Regularisierungsterm hinzu, der lokale Modelle dafür bestraft, sich zu weit vom globalen Modell zu entfernen. Der Hyperparameter mu steuert den Kompromiss: Höhere Werte halten die Clients näher am globalen Modell, können aber lokale Daten underfitten.

Privatsphäregarantien durch Differential Privacy

Federated Learning allein garantiert keine Privatsphäre. Modell-Updates können Informationen über Trainingsdaten durch Gradient-Inversion-Angriffe preisgeben. Differential Privacy liefert mathematische Garantien.

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

Der Privacy-Utility-Trade-off ist real. Mehr Rauschen bedeutet stärkere Privatsphäregarantien, aber langsamere Konvergenz. In der Praxis musst du das Privacy-Budget (Epsilon) über Runden hinweg verfolgen und das Training stoppen, wenn das Budget erschöpft ist.

Strategien zur Kommunikationseffizienz

Bandbreite ist der Flaschenhals beim Federated Learning. Vollständige Modell-Updates von tausenden Geräten zu senden, ist teuer. Kompressionstechniken reduzieren die Kommunikationskosten drastisch.

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

Top-k-Sparsifikation behält typischerweise über 90% der Modellqualität bei, während nur 1-10% der Parameter übertragen werden. Kombiniert mit Quantisierung (Reduzierung von float32 auf int8), lässt sich eine 100-fache Kompression mit minimalem Genauigkeitsverlust erreichen.

Orchestrierung des gesamten Trainingszyklus

Das Zusammenfügen erfordert eine sorgfältige Orchestrierung von Client-Auswahl, Training, Aggregation und Evaluation über mehrere Runden.

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

Produktive Federated-Learning-Systeme ergänzen Fehlertoleranz (Clients, die mitten in einer Runde ausfallen), sichere Aggregation (der Server sieht keine individuellen Updates) und asynchrone Updates (nicht auf langsame Clients warten). Jede Schicht erhöht die Komplexität, adressiert aber echte Herausforderungen beim Deployment.

Wichtige Erkenntnisse

Federated Learning bedeutet einen grundlegenden Wandel in der Art, wie wir über Daten und Modelltraining denken. Die zentrale Erkenntnis ist, dass Modelle aus Daten lernen können, die sie nie direkt sehen — und das ermöglicht Anwendungsfälle, die zuvor aufgrund von Datenschutzvorschriften, Wettbewerbsbedenken oder schlicht logistischer Einschränkungen unmöglich waren.

Die technischen Herausforderungen sind real: Nicht-IID-Datenverteilungen verschlechtern die Modellqualität, die Kommunikationskosten skalieren mit Modellgröße und Client-Anzahl, und Privatsphäregarantien gehen auf Kosten der Modellgenauigkeit. Aber das Feld reift schnell. Wenn deine Anwendung sensible Nutzerdaten involviert und du bisher vom zentralisierten Trainingsparadigma eingeschränkt warst, könnte Federated Learning der Ansatz sein, der deinen nächsten Durchbruch ermöglicht.

Wilfredo Rujel

Wilfredo Rujel

Full-Stack-Softwareentwickler

Diesen Beitrag teilenX