Skip to content

Customized Client Training Loops

SCAFFOLD

SCAFFOLD uses server and client control variates to correct local updates. This example implements corrected Option II with a composable trainer and processors that receive [weights, server_controls] and return [weights, delta_ci].

From the repository root:

uv sync
cd examples/customized_client_training
uv run python scaffold/scaffold.py -c scaffold/scaffold_MNIST_lenet5.toml

The shipped configuration downloads MNIST and trains LeNet-5 with five clients, selecting two per round. It keeps SGD with learning rate 0.01, momentum 0.9, and zero weight decay. Add --cpu for CPU execution or -b /path/to/new-run for a separate base directory for data and results.

Reference: Karimireddy et al., "SCAFFOLD: Stochastic Controlled Averaging for Federated Learning," in Proc. International Conference on Machine Learning (ICML), 2020; Algorithm 1 and equations (3)–(5).

Alignment with the paper

Let xx be the received model, yy the local model, cc the server control, cic_i the previous client control, and KK the number of completed optimizer updates. For vanilla SGD, the local update is y←y−η(g−ci+c)y \leftarrow y - \eta(g - c_i + c). The strategy applies −η(c−ci)-\eta(c-c_i) after the optimizer step, then computes corrected Option II: cinew=ci−c+(x−y)/(Kη)c_i^{\mathrm{new}} = c_i - c + (x-y)/(K\eta) and Δci=cinew−ci\Delta c_i = c_i^{\mathrm{new}} - c_i. These signs follow the main algorithm, not the reversed control terms in appendix equation (19). SCAFFOLDUpdateStrategyV2 remains a compatibility implementation of Option II; true Option I is not implemented.

The actual optimizer must use a finite, positive scalar learning rate, equal across participating parameter groups and constant within a local round. A different constant rate is allowed next round. Unequal rates or a change before a later update in the same round are rejected before that update. The correction and denominator use the rate actually executed, including the final partial gradient-accumulation window. That window counts once in KK; microbatches and skipped updates do not count as optimizer updates. With no completed updates, the client retains cic_i and emits zero delta.

The server adds ∑i∈SΔci/N\sum_{i\in S}\Delta c_i/N to cc, where \(N = \texttt{clients.total_clients}\). Control deltas are neither divided by the number of participants nor sample-weighted. Model aggregation does use sample weights, which differs from the paper's uniform-client model update when counts are unequal. The shipped momentum 0.9 and other non-vanilla optimizers also make the additive post-optimizer correction an extension. The paper's vanilla-SGD, uniform-client convergence guarantees are not claimed for these extensions.

The server validates received model and control payloads before staging, and checks the model/control state used after receive callbacks. Model and server controls commit only after aggregation callbacks and final validation succeed. Failures before that boundary discard staged controls and preserve the previous committed model/control values. Later evaluation or reporting failures retain the completed aggregation. External callback side effects are not rolled back.

The strategy persists each logical client's controls in scaffold_cv_<client_id>.pkl under the model directory (or its explicit save_path). Canonical state wins. If absent, a nonzero client can import the exact same-client <model_name>_<client_id>_control_variate.pth within that root, or the historical concatenated <root>scaffold_cv_<client_id>.pkl path, in that order. Invalid canonical state fails instead of falling back. Known historical buffer/frozen-parameter entries are ignored; required trainable controls must have matching shapes and finite floating-point values, with no unknown keys. Client 0 state is never used for another client. Distinct legacy files are preserved; subsequent saves use the canonical path.

Direct training accepts the result and persists client controls only after training callbacks, end hooks, and cleanup succeed. With trainer.max_concurrency (shipped as 2), a worker returns provisional controls and delta; the parent checks the current model/control handoff before accepting and saving controls. Failed end hooks, cleanup, or handoffs preserve previously accepted client controls and refuse an outbound delta. The child does not overwrite canonical controls. This state transfer and per-client persistence are not a complete server, optimizer, random-state, or privacy-state restart mechanism.

Corrected continuation changes the numerical evolution of structurally loadable old state. For comparisons, start a fresh run in a separate base directory with all server and client controls initialized consistently, initially to zero.


FedProx

To better handle system heterogeneity, the FedProx algorithm introduced a proximal term in the optimizer used by local training on the clients. It has been quite widely cited and compared with in the federated learning literature.

cd examples/customized_client_training
uv run fedprox/fedprox.py -c fedprox/fedprox_MNIST_lenet5.toml

Reference: T. Li, A. K. Sahu, M. Zaheer, M. Sanjabi, A. Talwalkar, V. Smith. "Federated Optimization in Heterogeneous Networks," in Proc. Machine Learning and Systems (MLSys), 2020.

Alignment with the paper

plato/trainers/strategies/algorithms/fedprox_strategy.py:111-193 snapshots the global iterate wtw^t at round start and augments the loss with (μ/2)∗∣∣w−wt∣∣(\mu / 2) * ||w - w^t||, which is the FedProx objective hk(w;wt)=Fk(w)+(mu/2)∗∣∣w−wt∣∣2h_k(w; w^t) = F_k(w) + (mu / 2) * ||w - w^t||^2 defined in Section 3 of Li et al. (2020). Autograd therefore produces the perturbed-gradient step without requiring a bespoke optimizer.

The config-aware wrapper FedProxLossStrategyFromConfig (plato/trainers/strategies/algorithms/fedprox_strategy.py:208-247) reads μ\mu from the same knobs (clients.proximal_term_penalty_constant / algorithm.fedprox_mu) that the paper exposes in Algorithms 1 and 2, so experiments reproduce the authors' hyperparameter schedules.

The reference TensorFlow release (litian96/FedProx/flearn/optimizer/pgd.py#L27-L92) applies an identical perturbation, computing g+μ∗(w−wt)g + \mu * (w - w^t) before the gradient step; Plato mirrors that logic in PyTorch by letting the proximal penalty backpropagate through the loss term, yielding a line-for-line correspondence with Perturbed Gradient Descent.


FedDyn

FedDyn couples a dynamic local regularizer with a dedicated server that combines selected client models and population history. Use its public entrypoint to install all three components: client, trainer, and server.

From the repository root:

uv sync
cd examples/customized_client_training
uv run python feddyn/feddyn.py -c feddyn/feddyn_MNIST_lenet5.toml \
  --cpu -b ./runtime/feddyn

The shipped configuration downloads MNIST and trains LeNet-5 with 1,000 clients, 10 selected per round, 20 local epochs, and up to three rounds. It uses alpha_coef = 0.01, uniform weighting, and SGD with learning rate 0.03, zero momentum, and zero weight decay. trainer.max_concurrency = 3 enables spawned local workers. Copy the TOML before changing settings and choose a separate base directory for independent runs.

Reference: Acar, D.A.E., Zhao, Y., Navarro, R.M., Mattina, M., Whatmough, P.N. and Saligrama, V. "Federated Learning Based on Dynamic Regularization," in Proc. International Conference on Learning Representations (ICLR), 2021; Algorithm 1 and the authors' displacement formulation.

Alignment with the paper

Let xx be the received cloud model, hih_i a client's cumulative displacement (initially zero), and yiy_i its local endpoint. The objective over trainable parameters is Ji(w)=Fi(w)+αi⟨hi,w⟩+(αi/2)∥w−x∥2J_i(w)=F_i(w)+\alpha_i\langle h_i,w\rangle+(\alpha_i/2)\|w-x\|^2, with gradient ∇Fi(w)+αi(w−x+hi)\nabla F_i(w)+\alpha_i(w-x+h_i). History is a cumulative displacement, not a measured gradient. For accepted participants SS, hinew=hi+yi−xh_i^{\mathrm{new}}=h_i+y_i-x; inactive clients retain their histories. The server computes xnew=∑i∈Syi/∣S∣+∑i=1Nhinew/Nx_{\mathrm{new}}=\sum_{i\in S}y_i/|S|+\sum_{i=1}^{N}h_i^{\mathrm{new}}/N. The population size NN is fixed, and inactive histories remain in the second mean. The corrected cloud model is evaluated and broadcast; the reference also reports separate selected-client and all-client averages.

With algorithm.feddyn_weighting = "uniform" (default), αi=α\alpha_i=\alpha. Sample mode requires algorithm.feddyn_sample_counts: positive integer counts for the full population, ordered by logical client ID from 1. It uses qi=Nni/∑jnjq_i=Nn_i/\sum_j n_j and αi=α/qi\alpha_i=\alpha/q_i. Neither mean becomes selected-client sample-weighted FedAvg; the retained algorithm.type = "fedavg" selects only underlying model exchange.

The regularizer is entirely in the loss, with zero optimizer weight decay. This matches the reference's regularizer gradient when clipping is inactive, but does not reproduce its active clipping order. No original convergence, accuracy, or communication-budget guarantee is claimed. The server stores all histories and sends the assigned history with each model.

Counts must match actual nonempty partition samplers, not label values, minibatch sizes, backing dataset size, or merely data.partition_size. Sample mode checks each configured count; uniform mode records accepted counts and requires them to remain stable. Uniform mode rejects a sample-count vector. algorithm.alpha_coef takes precedence over algorithm.feddyn_alpha, defaulting to 0.01. The full example requires positive finite alpha. Separately, FedDynLossStrategy(alpha=0) is task loss only, not a full-example FedAvg mode.

Qualification covers CPU float32/float64 execution. The supported boundary is synchronous participation with a fixed population, a fixed finite float32/float64 trainable set, and plain SGD at a fixed positive finite rate. Momentum, dampening, weight decay, Nesterov, maximization, schedulers, AMP, clipping, DP, asynchronous/cross-silo rounds, mutable buffers, and changed parameter ownership or schema are rejected. Gradient accumulation is supported, including a normalized final partial window; at least one optimizer update must complete.

The server owns all histories. Direct and spawned training return provisional results for the current run, round, and logical client. The parent validates the model/history handoff before allowing an outbound result. Workers do not save live history files or substitute one client's history for another's. The server commits a complete selected batch only after aggregation callbacks and final validation succeed. Earlier failures preserve committed model/history values; later evaluation or reporting failures retain the commit. External callback side effects are outside this guarantee.

Append --resume to the command above to resume from the last saved committed round in the same base directory. The shipped configuration writes models/feddyn/mnist/feddyn_lenet5.pth under that directory. Its atomic bundle contains the full model, all histories, counts, settings/schema, run and round identity, accepted dispatch tokens, and separate global Python, client-selection, NumPy, and Torch CPU RNG states. Resume validates the bundle and installs RNG states once before server registration and selection. Use compatible settings and the same partitions; trainer.rounds is the total desired round count. Model-only files cannot resume, and unfinished rounds, optimizer state, GPU RNG, and external state are not recovered.

Legacy history inspection is explicit and read-only through FedDynUpdateStrategy.read_legacy_history(context): the same client's <root>/feddyn_grad_<client_id>.pth precedes the exact old concatenated <root>_feddyn_grad_<client_id>.pth. Client 0 is not a substitute. Inspection validates tensors without adopting or rewriting them; save_path is only an inspection root. For a model-only warm start, call the dedicated server's warm_start_model(weights) after model initialization and before dispatch. It starts a new run with zero histories; there is no CLI warm-start flag. Older loss/aggregation trajectories do not implicitly continue under the corrected rules. Start fresh for comparisons.


MOON

MOON (Model-Contrastive Federated Learning) enhances standard FedAvg by adding a model-level contrastive regularizer. Each client augments the shared model with a projection head, clones the incoming global model as a positive anchor, and reuses a small buffer of its historical checkpoints as negatives. The server still performs sample-weighted averaging but records a short history of global states for downstream analysis or warm restarts.

cd examples/server_aggregation/moon/
uv run moon.py -c moon_MNIST_lenet5.toml

Key configuration parameters:

  • algorithm.mu: Weight assigned to the contrastive term (default: 5.0).
  • algorithm.temperature: Softmax temperature applied to cosine similarities (default: 0.5).
  • algorithm.history_size: Number of historical local models cached per client as negatives (default: 2).
  • trainer.model_name: Name used for checkpointing the projection-ready backbone (default: moon_lenet5).

Reference: Qinbin Li, Bingsheng He, Dawn Song. “Model-Contrastive Federated Learning,” in Proc. CVPR, 2021.

Alignment with the paper

Here’s how Plato's implementation lines up with Li et al. (CVPR 2021) and the authors’ reference implementation:

  • Projection head & representations – moon_model.py:31-79 implements the LeNet-style backbone plus a two-layer projection head, returning both logits and L2-normalised embeddings. The paper’s Eq. (3) (and typical contrastive-learning practice) calls for that projection step; the public repo’s simple CNN head even hints at it (they keep the projection MLP commented out). So keeping the projection in our model is faithful and helps the cosine similarities stay well behaved.

  • Local training objective – moon_trainer.py:26-152 combines the supervised cross-entropy with the temperature-scaled contrastive loss exactly like Eq. (1): positives come from the frozen global model, negatives from the stored local-history models, using the same μ\mu and τ\tau hyper-parameters exposed in the config (moon_MNIST_lenet5.toml:41-45). This mirrors train_net_fedcon in the reference implementation, which also weights the contrastive term by μ\mu and uses CrossEntropy on logits built from cosine similarities.

  • Historical model buffer – the client keeps a FIFO queue of past local checkpoints (moon_client.py:21-64), equivalent to model_buffer_size in the paper and the author's reference implementation; that buffer is fed into the trainer through the strategy context so MOON always has negatives available.

  • Server aggregation – the server still performs sample-weighted FedAvg (moon_server.py:12-35, moon_server_strategy.py:19-63), matching the MOON design which leaves the aggregation rule unchanged. The extra global-history deque is bookkeeping-only.

  • Shared architecture – moon.py:8-15 now instantiates MoonModel once and passes it into both the client and server (model=model). That guarantees the projection-enabled architecture is shared exactly, as required for the contrastive comparisons.

The only intentional deviation is that we L2-normalise the projection outputs before computing cosine similarities (moon_model.py:76-79), which the paper assumes implicitly and improves stability. Aside from that, the workflow, hyper-parameters, and loss all line up with the CVPR paper and the publicly released PyTorch reference.


FedMoS

FedMoS is a communication-efficient FL framework with coupled double momentum-based update and adaptive client selection, to jointly mitigate the intrinsic variance.

cd examples/customized_client_training
uv run fedmos/fedmos.py -c fedmos/fedmos_MNIST_lenet5.toml

Reference: X. Wang, Y. Chen, Y. Li, X. Liao, H. Jin and B. Li, "FedMoS: Taming Client Drift in Federated Learning with Double Momentum and Adaptive Selection," IEEE INFOCOM 2023.

Alignment with the paper

plato/trainers/strategies/algorithms/fedmos_strategy.py:104-205 implements FedMoS double-momentum update by first computing dt=gt+(1−a)∗(dt−1−gt−1)d_t = g_t + (1 - a) * (d_{t-1} - g_{t-1}) and then stepping w=(1−μ)∗w−η∗dt+μ∗wglobalw = (1 - \mu) * w - \eta * d_t + \mu * w_{\textrm{global}}; these are the same recursions described in Algorithm 1 of Wang et al. (2023).

The training loop enforces the paper's sequencing: FedMosStepStrategy.training_step (plato/trainers/strategies/algorithms/fedmos_strategy.py:487-538) calls update_momentum() immediately after backward() and passes the cached global model from FedMosUpdateStrategy.on_train_start (plato/trainers/strategies/algorithms/fedmos_strategy.py:329-347) into the optimizer step so the proximal pull uses the broadcast parameters from the server.

The official repository (Distributed-Learning-Networking-Group/FedMoS/optimizers/fedoptimizer.py#L27-L92) mirrors the same gradient-difference momentum and proximal correction, confirming the one-to-one correspondence between the Plato optimizer and the authors' release.