Implement a custom attack#
Problem: You want to stress-test an aggregator against a novel attack not yet in the library. How do you implement a custom Byzantine attack strategy from scratch?
All attacks in Krum follow the same protocol:
Subclass
AttackImplement
generate()as a@classmethodAccept
honest_gradients(first positional arg), an optionalouttensor, a keyword-onlyf(number of Byzantine gradients to produce), and any attack-specific keyword argumentsReturn a tensor of shape
(f, d)
What we’ll build#
We’ll implement RepeatAttack, the simplest possible Byzantine attack.
It takes the first honest gradient and repeats it f times. Every
Byzantine worker sends the exact same gradient. The intuition: if honest
workers are converging, sending a stale or misleading gradient repeated
many times can shift the aggregate away from the true direction.
Step 1: subclass Attack#
Create a file repeat_attack.py:
from collections.abc import Sequence
from typing import Any
from torch import Tensor
from krum.primitives.attacks import Attack
class RepeatAttack(Attack):
...
Step 2: implement generate (gradients only)#
Start without out or **specialized to see the core logic:
@classmethod
def generate(
cls,
honest_gradients: Sequence[Tensor] | Tensor,
*,
f: int,
) -> Tensor:
...
Stack the honest gradients, take the first one, and repeat it:
@classmethod
def generate(
cls,
honest_gradients: Sequence[Tensor] | Tensor,
*,
f: int,
) -> Tensor:
if not isinstance(honest_gradients, Tensor):
honest_gradients = stack(list(honest_gradients))
_, d = honest_gradients.shape
first = honest_gradients[0:1] # (1, d)
if f == 0:
return honest_gradients.new_empty((0, d))
return first.expand(f, d) # broadcast to (f, d)
If f is zero the attack returns an empty tensor. The simulation handles
this correctly.
Step 3: add the out parameter#
Add support for the optional output buffer:
@classmethod
def generate(
cls,
honest_gradients: Sequence[Tensor] | Tensor,
/,
out: Tensor | None = None,
*,
f: int,
) -> Tensor:
if not isinstance(honest_gradients, Tensor):
honest_gradients = stack(list(honest_gradients))
_, d = honest_gradients.shape
first = honest_gradients[0:1]
if f == 0:
empty = honest_gradients.new_empty((0, d))
if out is not None:
return out.copy_(empty)
return empty
result = first.expand(f, d)
if out is not None:
return out.copy_(result)
return result
Note
The out parameter is part of every attack’s signature. It
enables the caller to control memory allocation. See the
Attacks for the full API
reference.
Step 4: add specialized keyword arguments#
Add **specialized to absorb whatever the simulation passes:
@classmethod
def generate(
cls,
honest_gradients: Sequence[Tensor] | Tensor,
/,
out: Tensor | None = None,
*,
f: int,
**specialized: Any,
) -> Tensor:
...
To pass your own custom keyword arguments, use attack_kwargs on the
simulation:
sim = KrumSimulation(
...,
attack=RepeatAttack,
attack_kwargs={"my_param": 42},
)
Inside the attack, read them from specialized:
class RepeatAttack(Attack):
@classmethod
def generate(cls, honest_gradients, /, out=None, *, f, **specialized):
threshold = specialized.get("my_param", 0)
...
Full code#
from collections.abc import Sequence
from typing import Any
from torch import Tensor, stack
from krum.primitives.attacks import Attack
class RepeatAttack(Attack):
"""Byzantine attack that repeats the first honest gradient f times.
Every Byzantine worker sends the same gradient. Simple to
implement, useful as a baseline for testing aggregators.
"""
@classmethod
def generate(
cls,
honest_gradients: Sequence[Tensor] | Tensor,
/,
out: Tensor | None = None,
*,
f: int,
**specialized: Any,
) -> Tensor:
if not isinstance(honest_gradients, Tensor):
honest_gradients = stack(list(honest_gradients))
_, d = honest_gradients.shape
first = honest_gradients[0:1]
if f == 0:
empty = honest_gradients.new_empty((0, d))
if out is not None:
return out.copy_(empty)
return empty
result = first.expand(f, d)
if out is not None:
return out.copy_(result)
return result
Using your attack#
Import it and call it like any built-in attack. Attacks are stateless, so you pass the class itself, never an instance:
import torch
from repeat_attack import RepeatAttack
honest = torch.randn(8, 100)
byzantine = RepeatAttack.generate(honest, f=2)
print(byzantine.shape) # (2, 100)
print(torch.allclose(byzantine[0], byzantine[1])) # True, all identical
In a simulation#
Pass the class (not an instance) to a simulation:
from krum.primitives.aggregators.multikrum import MultiKrum
from krum.primitives.data_partitioners.iid import IidPartitioner
from krum.primitives.models.mlp import Krum2017MLPMnist
from krum.simulations.centralised.krum_nips_2017 import KrumSimulation
from torchvision import datasets, transforms
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)
# Partition the training set into one dataset per worker
train_datasets = IidPartitioner.partition(train_set, n=10, seed=42)
sim = KrumSimulation(
model_cls=Krum2017MLPMnist,
train_datasets=train_datasets,
test_set=test_set,
aggregator=MultiKrum,
attack=RepeatAttack,
n=10, f=2, rounds=50, batch_size=64, lr=0.01, seed=42,
)
sim.setup()
for _ in range(50):
sim.step()
loss, accuracy = sim.evaluate()
Testing#
A minimal test suite. Run with pytest repeat_attack.py:
import pytest
import torch
from repeat_attack import RepeatAttack
def test_generates_correct_shape():
honest = torch.randn(8, 100)
byzantine = RepeatAttack.generate(honest, f=2)
assert byzantine.shape == (2, 100)
def test_generates_no_byzantine_when_f_is_zero():
honest = torch.randn(8, 100)
byzantine = RepeatAttack.generate(honest, f=0)
assert byzantine.shape == (0, 100)
def test_all_byzantine_are_identical():
honest = torch.randn(8, 100)
byzantine = RepeatAttack.generate(honest, f=3)
assert torch.allclose(byzantine[0], byzantine[1])
assert torch.allclose(byzantine[0], byzantine[2])
def test_byzantine_matches_first_honest():
honest = torch.ones(4, 10)
byzantine = RepeatAttack.generate(honest, f=1)
assert torch.allclose(byzantine[0], honest[0])
Next steps#
Implement a custom aggregator: write an aggregation rule that defends against your attack.
Centralised simulation walkthrough: test your attack in a full training loop.
Decentralised simulation walkthrough: use the same attack in a peer-to-peer setting.
Structured experiments: benchmark your attack across seeds and aggregators with
Orchestrator.Browse the Attacks for all built-in attacks.