- Overview
- System Architecture
- Component Diagram
- Data Flow
- Communication Protocol
- Concurrency Model
- Design Patterns
- Scalability
FedLearn is a distributed federated learning framework built on gRPC that enables privacy-preserving machine learning across multiple clients without centralizing data. The framework supports both small models (CNNs) and large models (LLMs) through adaptive streaming mechanisms.
- Distributed Training: Multiple clients train independently on local data
- Model Agnostic: Works with any PyTorch model
- Adaptive Streaming: Automatically switches between unary and streaming based on model size
- Heartbeat Monitoring: Tracks client health in real-time
- Strategy Pattern: Pluggable aggregation strategies (FedAvg, custom)
- Large Model Support: Handles models up to several GB through chunked transfer
- Separation of Concerns: Clear boundaries between client, server, and communication layers
- Extensibility: Easy to add new strategies, clients, and protocols
- Robustness: Handles network failures, client dropouts, and stragglers
- Efficiency: Minimizes communication overhead and memory usage
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β FEDERATED LEARNING β
β ORCHESTRATION β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β
ββββββββββββββββββββΌβββββββββββββββββββ
β β β
βΌ βΌ βΌ
ββββββββββββββββ ββββββββββββββββ ββββββββββββββββ
β CLIENT 1 β β CLIENT 2 β β CLIENT N β
β β β β β β
β Local Data β β Local Data β β Local Data β
β Local Model β β Local Model β β Local Model β
ββββββββββββββββ ββββββββββββββββ ββββββββββββββββ
β β β
ββββββββββββββββββββΌβββββββββββββββββββ
β gRPC
βΌ
βββββββββββββββββββ
β SERVER/ β
β COORDINATOR β
β β
β Global Model β
β Aggregation β
β Evaluation β
βββββββββββββββββββ
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β PRESENTATION TIER β
β (Client Application Layer) β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ€
β β’ Client Implementation (fit(), get_parameters()) β
β β’ Training Loop β
β β’ Data Loading β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β
βΌ
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β COMMUNICATION TIER β
β (gRPC Layer) β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ€
β β’ GrpcClient / GrpcServicer β
β β’ Protocol Buffers (Serialization) β
β β’ Streaming (Chunked Transfer) β
β β’ Heartbeat (Keep-Alive) β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β
βΌ
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β BUSINESS LOGIC TIER β
β (Server Coordination & Strategy) β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ€
β β’ FLCoordinator (Round Management) β
β β’ Strategy (Aggregation Logic) β
β β’ Evaluation β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
βββββββββββββββββββββββββββββ SERVER SIDE βββββββββββββββββββββββββββββ
β β
β ββββββββββββββββββββ β
β β server.py β β
β β ββββββββββ β β
β β start_server() β β
β β β’ Creates β β
β β coordinator β β
β β β’ Starts gRPC β β
β β β’ Runs training β β
β β loop β β
β ββββββββββ¬ββββββββββ β
β β creates β
β βΌ β
β ββββββββββββββββββββββββββββββββββββββββββββββββ β
β β coordinator.py (FLCoordinator) β β
β β ββββββββββββββββββββββββββββ β β
β β STATE MANAGEMENT β β
β β β’ _global_model_params (OrderedDict) β β
β β β’ _client_updates_received (List) β β
β β β’ _registered_clients (Set) β β
β β β’ current_round (int) β β
β β β’ client_heartbeats (Dict) β β
β β β β
β β SYNCHRONIZATION β β
β β β’ _lock (threading.Lock) β β
β β β’ _round_complete_event (threading.Event) β β
β β β’ heartbeat_lock (Lock) β β
β β β β
β β METHODS β β
β β β’ start_round() β β
β β β’ wait_for_round_to_complete() β β
β β β’ submit_client_update() β β
β β β’ _trigger_aggregation_and_evaluation() β β
β β β’ update_client_heartbeat() β β
β β β’ is_client_alive() β β
β ββββββββββββββββ¬ββββββββββββββββββββββββββββββββ β
β β uses β
β βΌ β
β ββββββββββββββββββββββββββββββββββ β
β β strategy.py (Strategy) β β
β β ββββββββββββββββββββ β β
β β β’ initialize_parameters() β β
β β β’ aggregate_fit() β β
β β β’ evaluate() β β
β β β β
β β βββββββββββββββββββββββββββ β β
β β β FedAvg β β β
β β β β’ FedAvgAggregator β β β
β β β β’ Weighted averaging β β β
β β βββββββββββββββββββββββββββ β β
β ββββββββββββββββββββββββββββββββββ β
β β
β ββββββββββββββββββββββββββββββββββββββββββββ β
β β grpc_servicer.py (RPC Handlers) β β
β β ββββββββββββββββββββββββββββ β β
β β β’ RegisterClient() β β
β β β’ GetGlobalModelStream() β β
β β β’ SubmitModelUpdateStream() β β
β β β’ Heartbeat() β β
β β β’ GetServerStatus() β β
β ββββββββββββββββββββββββββββββββββββββββββββ β
β β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
βββββββββββββββββββββββββββββ CLIENT SIDE βββββββββββββββββββββββββββββ
β β
β ββββββββββββββββββββββββββββββββ β
β β client.py β β
β β βββββββββββ β β
β β Client (ABC) β β
β β β’ get_parameters() β β
β β β’ fit() β β
β β β β
β β start_client() β β
β β 1. Register β β
β β 2. Start heartbeat β β
β β 3. Main loop: β β
β β - Get global model β β
β β - Train (fit) β β
β β - Submit update β β
β β 4. Cleanup β β
β ββββββββββ¬ββββββββββββββββββββββ β
β β uses β
β βΌ β
β ββββββββββββββββββββββββββββββββββββββββββββ β
β β grpc_client.py (GrpcClient) β β
β β ββββββββββββββββββββββββββββ β β
β β TWO CHANNELS β β
β β β’ channel (main) β β
β β β’ heartbeat_channel β β
β β β β
β β METHODS β β
β β β’ register() β β
β β β’ get_global_model() β β
β β β’ submit_update() β β
β β ββ> _submit_update_unary() β β
β β ββ> _submit_update_stream() β β
β β β’ start_heartbeat() β β
β β β’ send_heartbeat() β β
β β β’ update_status() β β
β β β β
β β ADAPTIVE LOGIC β β
β β β’ STREAMING_THRESHOLD_MB = 100 β β
β β β’ Detects transformer models β β
β ββββββββββββββββββββββββββββββββββββββββββββ β
β β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
ββββββββββββββββββββββ COMMUNICATION LAYER βββββββββββββββββββββββββββββ
β β
β ββββββββββββββββββββββββββββββ β
β β fedlearn.proto β β
β β βββββββββββββ β β
β β MESSAGES β β
β β β’ Tensor β β
β β β’ ModelParameters β β
β β β’ ModelChunk β β
β β β’ ModelUpdateChunk β β
β β β’ HeartbeatRequest/ β β
β β Response β β
β β β β
β β SERVICE β β
β β β’ RegisterClient β β
β β β’ GetGlobalModelStream β β
β β β’ SubmitModelUpdateStream β β
β β β’ Heartbeat β β
β ββββββββββββββ¬ββββββββββββββββ β
β β generates β
β βΌ β
β βββββββββββββββββββββββββββββ β
β β generated/ β β
β β β’ fedlearn_pb2.py β β
β β β’ fedlearn_pb2_grpc.py β β
β βββββββββββββββββββββββββββββ β
β β
β ββββββββββββββββββββββββββββββββββββββββββ β
β β serializer.py β β
β β ββββββββββββββ β β
β β SMALL MODELS (Unary) β β
β β β’ parameters_to_proto() β β
β β β’ proto_to_parameters() β β
β β β β
β β LARGE MODELS (Streaming) β β
β β β’ parameters_to_chunks() β β
β β - Serialize with torch.save β β
β β - Optional LZ4 compression β β
β β - Split into 50MB chunks β β
β β β’ chunks_to_parameters() β β
β β - Concatenate chunks β β
β β - Decompress if needed β β
β β - Load with torch.load β β
β ββββββββββββββββββββββββββββββββββββββββββ β
β β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
INITIALIZATION PHASE
ββββββββββββββββββββ
Server:
1. start_server() called
2. Create FLCoordinator with Strategy
3. Set initial_parameters from strategy
4. Start gRPC server on specified port
5. Enter training loop
Clients (parallel):
1. start_client() called
2. Create GrpcClient
3. Register with server
4. Start heartbeat thread
5. Enter main loop
ROUND N BEGINS
ββββββββββββββ
Server (Main Thread):
ββββββββββββββββββββββββββββββββββββββββ
β coordinator.start_round() β
β ββ> _round_complete_event.clear() β
β β
β coordinator.wait_for_round_to_ β
β complete() β
β ββ> BLOCKS waiting for event β
ββββββββββββββββββββββββββββββββββββββββ
SERVER STATUS: WAITING
PHASE 1: MODEL DOWNLOAD
ββββββββββββββββββββββββ
Client 1:
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β params, round, config = β
β grpc_client.get_global_model() β
β β
β GrpcClient: β
β req = GetGlobalModelRequest(client_id) β
β for chunk in stub.GetGlobalModelStream(req): β
β chunks.append(chunk.chunk_data) β
β β
β full_data = b''.join(chunks) β
β with BytesIO(full_data) as buffer: β
β model_data = torch.load(buffer, weights_only=True)β
β return model_data['parameters'], round, config β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β²
β gRPC Stream
β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β Server (RPC Thread): β
β β
β GrpcServicer.GetGlobalModelStream(): β
β params = coordinator.get_global_model_for_client()β
β β
β buffer = BytesIO() β
β torch.save({'parameters': params}, buffer) β
β data = buffer.getvalue() β
β β
β for i in range(num_chunks): β
β yield ModelChunk( β
β chunk_index=i, β
β chunk_data=data[start:end] β
β ) β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
Timeline:
T=0s: Client requests model
T=1s: Server starts streaming chunk 1/20
T=3s: Chunk 5/20 transferred
T=5s: Chunk 10/20 transferred
T=8s: Chunk 20/20 transferred (final)
T=9s: Client reconstructs model
PHASE 2: LOCAL TRAINING
ββββββββββββββββββββββββ
Client 1:
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β new_params, num_examples = client.fit(params, config)β
β β
β MyClient.fit(): β
β model.load_state_dict(params) β
β model.train() β
β β
β for epoch in range(5): β
β for batch_idx, (data, target) in train_loader: β
β # Training step β
β loss.backward() β
β optimizer.step() β
β β
β # Update status every 10 batches β
β if batch_idx % 10 == 0: β
β grpc_client.update_status( β
β "training", batch_idx, total_batches β
β ) β
β β
β return model.state_dict(), len(dataset) β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
Parallel (Background Thread):
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β Heartbeat Thread (every 5 seconds): β
β grpc_client.send_heartbeat() β
β ββ> HeartbeatRequest( β
β client_id="client_1", β
β status="training", β
β current_step=45, β
β total_steps=100, β
β current_round=1 β
β ) β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β
βΌ gRPC (separate channel)
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β Server (RPC Thread): β
β GrpcServicer.Heartbeat(): β
β coordinator.update_client_heartbeat(...) β
β ββ> Update timestamp, print progress β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
Timeline:
T=10s: Training starts
T=15s: Heartbeat sent (Step 0/100)
T=20s: Heartbeat sent (Step 20/100)
T=25s: Heartbeat sent (Step 40/100)
...
T=60s: Training complete
PHASE 3: MODEL UPLOAD
ββββββββββββββββββββββ
Client 1:
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β success = grpc_client.submit_update( β
β new_params, num_examples, round β
β ) β
β β
β GrpcClient.submit_update(): β
β # Calculate size β
β size_mb = calculate_size(params) β
β β
β # Decide: streaming or unary? β
β if size_mb > 100 or is_transformer: β
β return _submit_update_stream(...) β
β β
β GrpcClient._submit_update_stream(): β
β def chunk_generator(): β
β for chunk_info in parameters_to_chunks(params): β
β yield ModelUpdateChunk( β
β chunk_data=chunk_info['chunk_data'] β
β ) β
β β
β response = stub.SubmitModelUpdateStream( β
β chunk_generator() β
β ) β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β
βΌ gRPC Stream
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β Server (RPC Thread): β
β β
β GrpcServicer.SubmitModelUpdateStream(): β
β chunks = [] β
β for chunk in request_iterator: β
β chunks.append(chunk.chunk_data) β
β β
β full_data = b''.join(chunks) β
β params, num_examples = β
β chunks_to_parameters(full_data) β
β β
β coordinator.submit_client_update( β
β client_id, params, num_examples, round β
β ) β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β
βΌ
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β Coordinator.submit_client_update(): β
β with self._lock: β
β # Validate round number β
β if trained_on_round != self.current_round: β
β return # Ignore stale/future updates β
β β
β # Add to list β
β self._client_updates_received.append( β
β (params, num_examples) β
β ) β
β β
β # Check if we have enough β
β if len(self._client_updates_received) == β
β self.clients_per_round: β
β β
β self._trigger_aggregation_and_evaluation() β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
Timeline:
T=61s: Client 1 starts upload
T=65s: Chunk 5/18 uploaded
T=70s: Chunk 10/18 uploaded
T=75s: Chunk 18/18 uploaded
T=76s: Server reconstructs model
T=76s: Coordinator receives update (1/3 clients)
[Clients 2 and 3 repeat phases 1-3 in parallel]
T=120s: Client 2 submits update (2/3 clients)
T=180s: Client 3 submits update (3/3 clients)
T=180s: TRIGGER AGGREGATION
PHASE 4: AGGREGATION
ββββββββββββββββββββ
Coordinator (Main Thread Wakes Up):
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β _trigger_aggregation_and_evaluation(): β
β β
β # Get all updates β
β results = list(self._client_updates_received) β
β # [(params_1, 1000), (params_2, 500), β
β # (params_3, 1500)] β
β β
β self._client_updates_received.clear() β
β β
β # Aggregate β
β aggregated_params = self.strategy.aggregate_fit( β
β self.current_round, β
β results β
β ) β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β
βΌ
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β Strategy.aggregate_fit(): β
β return self.aggregator.aggregate(results) β
β β
β FedAvgAggregator.aggregate(): β
β total_examples = 1000 + 500 + 1500 = 3000 β
β β
β aggregated = OrderedDict() β
β for params, num_examples in results: β
β weight = num_examples / total_examples β
β # 0.333, 0.167, 0.500 β
β β
β for key in params: β
β torch.add(aggregated[key], params[key], β
β alpha=weight, out=aggregated[key]) β
β β
β return aggregated β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
Mathematical Example:
Client 1 weight: layer.weight = [[1.0, 2.0], [3.0, 4.0]] (1000 samples)
Client 2 weight: layer.weight = [[2.0, 3.0], [4.0, 5.0]] (500 samples)
Client 3 weight: layer.weight = [[1.5, 2.5], [3.5, 4.5]] (1500 samples)
Aggregated = 0.333*[[1.0, 2.0], [3.0, 4.0]] +
0.167*[[2.0, 3.0], [4.0, 5.0]] +
0.500*[[1.5, 2.5], [3.5, 4.5]]
= [[1.417, 2.417], [3.417, 4.417]]
PHASE 5: EVALUATION
βββββββββββββββββββ
Coordinator:
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β # Update global model β
β self._global_model_params = aggregated_params β
β β
β # Evaluate β
β loss, metrics = self.strategy.evaluate( β
β self.current_round, β
β self._global_model_params β
β ) β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β
βΌ
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β Strategy.evaluate(): β
β if self.evaluate_fn is None: β
β return None β
β β
β # User-provided evaluation function β
β loss, metrics = self.evaluate_fn( β
β server_round, β
β parameters β
β ) β
β β
β # Example evaluate_fn: β
β def evaluate_fn(round_num, params): β
β model.load_state_dict(params) β
β model.eval() β
β β
β total_loss = 0 β
β correct = 0 β
β for data, target in test_loader: β
β output = model(data) β
β loss = criterion(output, target) β
β total_loss += loss.item() β
β correct += (output.argmax(1) == target).sum() β
β β
β accuracy = correct / len(test_loader.dataset) β
β return total_loss / len(test_loader), { β
β 'accuracy': accuracy β
β } β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
Timeline:
T=181s: Aggregation complete
T=182s: Evaluation starts
T=185s: Evaluation complete
Output:
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β [Server] Round 1 complete. β
β Metrics: {'loss': 1.2345, 'accuracy': 0.78} β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
PHASE 6: ROUND COMPLETION
ββββββββββββββββββββββββββ
Coordinator:
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β self.latest_metrics = {"loss": loss, **metrics} β
β β
β # Signal round complete β
β self._round_complete_event.set() β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β
βΌ UNBLOCKS
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β Server (Main Thread): β
β # wait_for_round_to_complete() returns β
β β
β metrics = coordinator.get_latest_metrics() β
β history.append((round_num, metrics)) β
β β
β # Advance to next round β
β coordinator.current_round += 1 β
β β
β # Loop continues for next round... β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
CLIENTS CONTINUE
ββββββββββββββββ
All Clients (in parallel):
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β while True: β
β params, server_round, config = β
β comm_client.get_global_model() β
β β
β if server_round == -1: β
β break # Training complete β
β β
β if server_round > last_completed_round: β
β # New round! Repeat phases 1-3 β
β ... β
β else: β
β # Server still on same round, wait β
β time.sleep(5) β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
TRAINING COMPLETION
βββββββββββββββββββ
Server (after num_rounds):
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β final_params = coordinator.get_global_model_params() β
β return history, final_params β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
Clients (detect completion):
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β params, server_round, config = β
β comm_client.get_global_model() β
β β
β if server_round == -1: β
β print("Server finished training. Shutting down.") β
β break β
β β
β # Cleanup β
β comm_client.stop_heartbeat() β
β comm_client.close() β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
service FederatedLearningService {
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β RegisterClient β
β Request: RegisterClientRequest β
β Response: RegisterClientResponse β
β Type: Unary β
β Purpose: Client registration at startup β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β GetGlobalModelStream β
β Request: GetGlobalModelRequest β
β Response: stream ModelChunk β
β Type: Server Streaming β
β Purpose: Download global model in chunks β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β SubmitModelUpdateStream β
β Request: stream ModelUpdateChunk β
β Response: SubmitModelUpdateResponse β
β Type: Client Streaming β
β Purpose: Upload trained model in chunks β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β Heartbeat β
β Request: HeartbeatRequest β
β Response: HeartbeatResponse β
β Type: Unary β
β Purpose: Keep-alive and status updates β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
}
Small Messages (<1KB):
β’ RegisterClientRequest
β’ RegisterClientResponse
β’ HeartbeatRequest
β’ HeartbeatResponse
β’ GetGlobalModelRequest
Large Messages (Variable):
β’ ModelChunk: 50 MB per chunk
β’ ModelUpdateChunk: 50 MB per chunk
Total Transfer Sizes (Example):
β’ Small CNN (10 MB): 1 chunk each direction
β’ Medium ResNet (150 MB): 3 chunks each direction
β’ GPT-2 (500 MB): 10 chunks each direction
β’ LLaMA-7B (14 GB): 280 chunks each direction
ββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β ADAPTIVE STREAMING DECISION β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββ€
β β
β if is_transformer(model): β
β USE STREAMING (reason: transformer detected) β
β β
β elif model_size_mb > 100: β
β USE STREAMING (reason: size > threshold) β
β β
β else: β
β USE UNARY (reason: small model) β
β β
ββββββββββββββββββββββββββββββββββββββββββββββββββββββ
Transformer Detection:
Keywords in layer names:
β’ 'transformer'
β’ 'bert'
β’ 'gpt'
β’ 'attention'
β’ 'encoder'
β’ 'decoder'
ββββββββββββββββββ SERVER ββββββββββββββββββ
β β
β Main Thread β
β βββββββββββββββββββββββββββββββββββββββ β
β β Training Loop β β
β β for round in rounds: β β
β β coordinator.start_round() β β
β β coordinator.wait_for_round_to_ β β
β β complete() [BLOCKS] β β
β β coordinator.current_round += 1 β β
β βββββββββββββββββββββββββββββββββββββββ β
β β
β gRPC ThreadPool (10 workers) β
β βββββββββββββββββββββββββββββββββββββββ β
β β Thread 1: RegisterClient RPC β β
β βββββββββββββββββββββββββββββββββββββββ€ β
β β Thread 2: GetGlobalModelStream RPC β β
β βββββββββββββββββββββββββββββββββββββββ€ β
β β Thread 3: SubmitModelUpdateStream β β
β βββββββββββββββββββββββββββββββββββββββ€ β
β β Thread 4: Heartbeat RPC β β
β βββββββββββββββββββββββββββββββββββββββ€ β
β β Thread 5: Heartbeat RPC β β
β βββββββββββββββββββββββββββββββββββββββ β
β β
β Synchronization β
β β’ coordinator._lock (state mutations) β
β β’ coordinator._round_complete_event β
β β’ coordinator.heartbeat_lock β
β β
βββββββββββββββββββββββββββββββββββββββββββββ
ββββββββββββββββββ CLIENT ββββββββββββββββββ
β β
β Main Thread β
β βββββββββββββββββββββββββββββββββββββββ β
β β Client Loop β β
β β while True: β β
β β get_global_model() β β
β β fit() β β
β β submit_update() β β
β βββββββββββββββββββββββββββββββββββββββ β
β β
β Heartbeat Thread (daemon) β
β βββββββββββββββββββββββββββββββββββββββ β
β β while heartbeat_active: β β
β β send_heartbeat() β β
β β sleep(5) β β
β βββββββββββββββββββββββββββββββββββββββ β
β β
β Two gRPC Channels β
β β’ Main channel (model transfer) β
β β’ Heartbeat channel (keep-alive) β
β β
βββββββββββββββββββββββββββββββββββββββββββββ
# Coordinator uses threading primitives
1. Lock for State Mutations
self._lock = threading.Lock()
with self._lock:
self._client_updates_received.append(update)
if len(self._client_updates_received) == threshold:
self._trigger_aggregation()
2. Event for Round Completion
self._round_complete_event = threading.Event()
# Server main thread
coordinator.wait_for_round_to_complete()
# β Blocks on event.wait()
# RPC thread (when aggregation done)
self._round_complete_event.set()
# β Wakes up main thread
3. Separate Lock for Heartbeats
self.heartbeat_lock = Lock()
# Prevents heartbeat updates from blocking
# client update submissions# Abstract strategy
class Strategy(ABC):
@abstractmethod
def aggregate_fit(self, results):
pass
# Concrete strategies
class FedAvg(Strategy):
def aggregate_fit(self, results):
return weighted_average(results)
class FedProx(Strategy):
def aggregate_fit(self, results):
return weighted_average_with_proximal_term(results)
# Usage
strategy = FedAvg(...) # Can swap with FedProx
coordinator = FLCoordinator(strategy=strategy)Benefits:
- Easy to add new aggregation algorithms
- Algorithms are interchangeable
- Clean separation of concerns
# Abstract template
class Client(ABC):
@abstractmethod
def fit(self, parameters, config):
"""Subclasses implement training logic"""
pass
@abstractmethod
def get_parameters(self):
"""Subclasses implement parameter extraction"""
pass
# Concrete implementation
class CNNClient(Client):
def fit(self, parameters, config):
# Custom training for CNN
pass
def get_parameters(self):
return self.model.state_dict()
# Common behavior in start_client()
def start_client(client):
while True:
params = get_model()
new_params = client.fit(params) # Calls subclass method
submit(new_params)Benefits:
- Common client lifecycle in one place
- Custom training logic in subclasses
- Prevents code duplication
# Complex gRPC internals
class GrpcClient:
def submit_update(self, params, num_examples, round):
# Hides complexity:
# - Size calculation
# - Streaming vs unary decision
# - Serialization
# - Chunking
# - Error handling
if should_stream(params):
return self._submit_update_stream(...)
else:
return self._submit_update_unary(...)
# Simple client interface
client = GrpcClient(client_id, address)
client.submit_update(params, 1000, 1) # Simple call!Benefits:
- Hides gRPC complexity
- Simple API for clients
- Easy to modify internals
# Coordinator observes client state
class FLCoordinator:
def update_client_heartbeat(self, client_id, status, ...):
self.client_heartbeats[client_id] = {
'status': status,
'last_seen': time.time()
}
# Print progress if needed
if current_step % 10 == 0:
print(f"Client {client_id}: {status}")
# Clients notify coordinator
class GrpcClient:
def _heartbeat_loop(self):
while active:
self.send_heartbeat() # Notifies coordinator
time.sleep(5)Benefits:
- Loose coupling
- Real-time monitoring
- No polling needed
# Generator yields chunks one at a time
def parameters_to_chunks(params) -> Generator[Dict, None, None]:
for i in range(num_chunks):
yield {
'chunk_index': i,
'chunk_data': data[start:end],
'is_final_chunk': (i == num_chunks - 1)
}
# Used in streaming
def chunk_generator():
for chunk_info in parameters_to_chunks(params):
yield ModelUpdateChunk(**chunk_info)
response = stub.SubmitModelUpdateStream(chunk_generator())Benefits:
- Memory efficient (one chunk at a time)
- Works with gRPC streaming
- Clean separation of chunking logic
NUMBER OF CLIENTS
βββββββββββββββββ
Current: 2-10 clients (tested)
Supported: Up to 100+ clients
Bottleneck: Server aggregation time
Solution:
β’ Asynchronous aggregation
β’ Hierarchical aggregation
β’ Client sampling
MODEL SIZE
ββββββββββ
Small: <100 MB (unary transfer)
Medium: 100 MB - 1 GB (streaming)
Large: 1 GB - 10 GB+ (streaming)
Tested: Up to 2 GB models
Theoretical: Limited by disk space and memory
COMMUNICATION EFFICIENCY
ββββββββββββββββββββββββ
Baseline: Full model transfer each round
Optimizations:
β’ Compression (LZ4): 2-3x reduction
β’ Gradient-only transfer: 1x size
β’ Sparse updates: Variable reduction
β’ Quantization: 4-8x reduction
COMPUTATION TIME
ββββββββββββββββ
Factors:
β’ Model size
β’ Dataset size
β’ Hardware (CPU/GPU)
β’ Number of local epochs
Example (CNN on MNIST):
β’ Local training: 2 minutes
β’ Model download: 5 seconds
β’ Model upload: 5 seconds
β’ Total per round: ~2.5 minutes
Example (GPT-2 on text):
β’ Local training: 30 minutes
β’ Model download: 45 seconds
β’ Model upload: 60 seconds
β’ Total per round: ~32 minutes
1. ADAPTIVE STREAMING
β’ Small models: Unary (faster)
β’ Large models: Streaming (reliable)
2. DUAL CHANNELS
β’ Main: Model transfer (blocks)
β’ Heartbeat: Keep-alive (non-blocking)
3. COMPRESSION
β’ Optional LZ4 compression
β’ 2-3x size reduction
β’ Minimal CPU overhead
4. THREADING
β’ 10 gRPC worker threads
β’ Concurrent client handling
β’ Background heartbeat threads
5. MEMORY EFFICIENCY
β’ Streaming serialization
β’ Chunk-by-chunk processing
β’ Immediate garbage collection
1. ASYNCHRONOUS AGGREGATION
Current: Wait for all clients
Future: Aggregate as clients arrive
2. CLIENT SELECTION
Current: Use all connected clients
Future: Sample subset per round
3. GRADIENT COMPRESSION
Current: Full parameter transfer
Future: Gradient sparsification
4. HIERARCHICAL AGGREGATION
Current: Flat client-server
Future: Multi-tier aggregation
5. DIFFERENTIAL PRIVACY
Current: No privacy guarantees
Future: DP-SGD, secure aggregation
βββββββββββββββββββββββββββββββββββ
β Single Machine β
βββββββββββββββββββββββββββββββββββ€
β Server Process β
β β’ Port 50051 β
β β
β Client Process 1 β
β Client Process 2 β
β Client Process 3 β
βββββββββββββββββββββββββββββββββββ
Use Case: Development, debugging
Command:
Terminal 1: python run_server.py
Terminal 2: python run_client.py --client-id=1
Terminal 3: python run_client.py --client-id=2
βββββββββββββββββββββββ
β Server Machine β
β IP: 192.168.1.100 β
β Port: 50051 β
βββββββββββββββββββββββ
β
β Internet/LAN
β
ββββββ΄βββββ¬βββββββββ¬βββββββββ
β β β β
βΌ βΌ βΌ βΌ
βββββββββ βββββββββ βββββββββ βββββββββ
βClient1β βClient2β βClient3β βClientNβ
β GPU β β GPU β β CPU β β GPU β
βββββββββ βββββββββ βββββββββ βββββββββ
Use Case: Real federated learning
Setup:
Server: python run_server.py --address=0.0.0.0:50051
Clients: python run_client.py --server=192.168.1.100:50051
ββββββββββββββββββββββββββββββββββββββββββββ
β AWS Cloud β
ββββββββββββββββββββββββββββββββββββββββββββ€
β EC2 Instance (Server) β
β β’ Type: t3.large β
β β’ Public IP: X.X.X.X β
β β’ Security Group: Port 50051 open β
β β
β EC2 Instances (Clients) β
β β’ Type: g4dn.xlarge (GPU) β
β β’ Private IPs: Connect to server β
ββββββββββββββββββββββββββββββββββββββββββββ
Docker Deployment:
docker run -p 50051:50051 fedlearn-server
docker run fedlearn-client --server=X.X.X.X:50051
-
gRPC for Communication
- Efficient binary protocol
- Built-in streaming support
- Cross-platform compatibility
-
Adaptive Streaming
- Handles models from 1MB to 10GB+
- Automatic decision based on size
- Robust to network issues
-
Dual-Channel Design
- Prevents heartbeat blocking
- Enables long-running transfers
- Improves reliability
-
Strategy Pattern for Aggregation
- Easy to extend
- Pluggable algorithms
- Clean separation of concerns
-
Thread-Safe Coordination
- Supports concurrent clients
- Prevents race conditions
- Enables synchronization
CHOSEN vs ALTERNATIVE
ββββββββββββββββββββββββββββββββββββββββββββββββββ
gRPC/Protobuf vs HTTP/REST + JSON
β’ Faster vs β’ Easier debugging
β’ Streaming vs β’ Better tooling
β’ Type-safe vs β’ More familiar
Synchronous Rounds vs Asynchronous
β’ Simpler logic vs β’ Faster convergence
β’ Better convergence vs β’ Handles stragglers
β’ Easier debugging vs β’ More efficient
Weighted Averaging vs Median/Krum
β’ Faster vs β’ Byzantine robust
β’ Standard vs β’ More complex
β’ Well-understood vs β’ Slower
This architecture provides a solid foundation for federated learning while remaining extensible for future enhancements.