Decentralised simulation walkthrough#
Problem: Your scenario needs peer-to-peer communication with per-worker models instead of a central parameter server. Each worker trains its own model and the simulation handles model mixing between neighbours each round.
Krum ships with one built-in decentralised (peer-to-peer) simulation. This tutorial covers the peer-to-peer framework where each worker holds its own model and workers exchange models through a communication topology.
All decentralised simulations share the lifecycle:
instantiate → step (or run(rounds)), with a per-round snapshot.
See also
- Decentralised simulations
Full reference for
MonnaSimulation.
Each round runs two phases:
Local optimisation: each honest worker computes a gradient on its own batch and updates its own model.
Model mixing: each worker gathers
n - fmodels from other nodes (received set) and replaces its model with an aggregate of that set.
Minimal example#
Data preparation#
The train_datasets argument of a decentralised simulation is a
sequence of :class:`~torch.utils.data.Dataset` instances, one per
worker. The simulation wraps each honest worker’s dataset into its own
DataLoader (batch size, shuffling) and
re-iterates it automatically once an epoch is exhausted, so a stream
never runs out:
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from krum.primitives.data_partitioners.iid import IidPartitioner
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
seed = 42
torch.manual_seed(seed)
n, f = 6, 0
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,)),
])
train_set = datasets.MNIST(root="./data", train=True, download=True, transform=transform)
test_set = datasets.MNIST(root="./data", train=False, download=True, transform=transform)
# One dataset per worker: IidPartitioner gives each worker an equal,
# shuffled portion of the training set
train_datasets = IidPartitioner.partition(train_set, n=n, seed=seed)
len(train_datasets) must equal n, including Byzantine workers:
only the first n - f are ever trained on (Byzantine workers craft
their models from the honest ones via the attack), but the full
n-length sequence is required.
For non-IID data,
DirichletPartitioner
produces per-class label skew (controlled by an alpha parameter).
See Working with Data Partitioners for the full partitioner
family; the experiment scripts in experiments/decentralised/ use
IidPartitioner and
DirichletPartitioner
directly.
Instantiation#
The model is wrapped in a Model container
that exposes a .parameters tensor and a .module (the underlying
nn.Module). Pass the model, the per-worker datasets, the test set,
and the hyperparameters to the simulation constructor:
from krum.primitives.models.mlp import Monna2023SmallMnist
from krum.primitives.models import Model
from krum.simulations.decentralised.monna_icml_2023 import MonnaSimulation
model = Model(Monna2023SmallMnist().to(device))
sim = MonnaSimulation(
model=model,
train_datasets=train_datasets,
train_batch_size=64,
test_set=test_set,
test_batch_size=256,
loss_fn=nn.CrossEntropyLoss(),
n=n,
f=f,
learning_rate=0.1,
beta=0.99,
seed=seed,
)
MonnaSimulation also accepts weight_decay (L2 regularization on
the honest gradients, 0.0 by default). Pass weight_decay=1e-4
to match the standard SGD-with-momentum recipe.
Training#
Call run(rounds) to train. The result is a list of result dicts, one
per round:
results = sim.run(50)
# results is a list of MonnaStepResult dicts, one per round
print(f"Ran {len(results)} rounds")
print(f"Final per-worker losses: {results[-1]['losses']}")
Inspecting the round snapshot#
step() returns a
StepResult dict (or a subclass like
MonnaStepResult):
result = sim.step() # a single round
The returned dict contains:
Key |
Description |
Shape |
|---|---|---|
|
Round counter (1-indexed) |
scalar |
|
Committed parameters after mixing |
|
|
Momentum buffer ( |
|
|
Computed gradients before local update |
|
|
Parameters after local update, before mixing |
|
|
Byzantine models injected this round |
|
|
Same as |
|
|
Per-worker scalar losses |
|
Byzantine workers#
Add Byzantine workers by setting f > 0 and providing an attack. Each
round the attack generates f Byzantine parameter vectors from the
honest ones. The
byzantine_reach
mode controls which workers receive them:
"all": every Byzantine model reaches every worker (worst-case adversary)."sampled": responders are drawn uniformly from all other nodes, so a worker receives0tofByzantine models.
from krum.primitives.attacks.sign_flip import SignFlipAttack
byzantine_datasets = IidPartitioner.partition(train_set, n=8, seed=seed)
sim_all = MonnaSimulation(
model=model, train_datasets=byzantine_datasets, train_batch_size=64,
test_set=test_set, test_batch_size=256, loss_fn=nn.CrossEntropyLoss(),
n=8, f=2, learning_rate=0.1,
attack=SignFlipAttack, attack_kwargs={"scale": 1.5},
byzantine_reach="all", seed=42,
)
sim_sampled = MonnaSimulation(
model=model, train_datasets=byzantine_datasets, train_batch_size=64,
test_set=test_set, test_batch_size=256, loss_fn=nn.CrossEntropyLoss(),
n=8, f=2, learning_rate=0.1,
attack=SignFlipAttack, attack_kwargs={"scale": 1.5},
byzantine_reach="sampled", seed=42,
)
result_all = sim_all.run(50)
result_sampled = sim_sampled.run(50)
print(f"'all' mean final loss: {result_all[-1]['losses'].mean():.4f}")
print(f"'sampled' mean final loss: {result_sampled[-1]['losses'].mean():.4f}")
Switching the mixing aggregator#
By default MonnaSimulation uses
NearestNeighborAverage
with num_closest = n - 2f. Override it with any
Aggregator subclass. Pass extra
aggregator parameters through aggregator_kwargs:
from krum.primitives.aggregators.median import Median
sim = MonnaSimulation(
model=model,
train_datasets=byzantine_datasets,
train_batch_size=64,
test_set=test_set,
test_batch_size=256,
loss_fn=nn.CrossEntropyLoss(),
n=8,
f=2,
learning_rate=0.1,
attack=SignFlipAttack,
aggregator=Median,
seed=42,
)
results = sim.run(50)
Evaluating worker models#
In the decentralised setting each worker has its own parameters. Evaluate every honest worker on the same test set and average the results. The function below loads each worker’s parameter vector into the shared model, runs the full test set, and averages across workers:
@torch.no_grad()
def evaluate_workers(model, parameters, test_loader, loss_fn):
losses, accuracies = [], []
for worker_params in parameters:
# copy worker params into the model
model.parameters.copy_(worker_params)
model.module.eval()
# run the full test set for this worker
total_loss = total_correct = total = 0
for inputs, targets in test_loader:
inputs, targets = inputs.to(device), targets.to(device)
logits = model.module(inputs)
loss = loss_fn(logits, targets)
total_loss += loss.item() * targets.numel()
total_correct += (logits.argmax(1) == targets).sum().item()
total += targets.numel()
losses.append(total_loss / total)
accuracies.append(total_correct / total)
# average across all honest workers
return sum(losses) / len(losses), sum(accuracies) / len(accuracies)
test_loader = DataLoader(test_set, batch_size=256, shuffle=False)
avg_loss, avg_acc = evaluate_workers(
model, sim.parameters, test_loader, nn.CrossEntropyLoss()
)
print(f"Average test loss: {avg_loss:.4f}, accuracy: {avg_acc:.2%}")
Next steps#
Centralised simulation walkthrough: try the simpler parameter-server simulation first if you haven’t.
Implement a custom aggregator: write your own aggregation rule and test it in this simulation.
Implement a custom attack: write your own Byzantine attack and test it in this simulation.
Structured experiments: collect structured results across multiple configurations with
MetricandOrchestrator.Decentralised simulations: the full decentralised simulation reference.