Skip to content

Repository files navigation

fl-secure-aggregation-zkp

Federated learning on MNIST with secure aggregation (Bonawitz-style pairwise masking), a continuous model-poisoning attack that exploits the privacy that secure aggregation buys, and a zero-knowledge norm-bound defence that restores robustness without breaking the privacy guarantee.

The whole stack is built from primitives in pure PyTorch + Python, so each layer is inspectable as a few hundred lines of code. The ZKP stage ships two interchangeable backends — a from-scratch Python NIZK (Pedersen + Schnorr-OR + Fiat–Shamir) used by default, and an optional Rust

  • arkworks Groth16 SNARK extension (03_zkp_defence/zkp_rust/) selected at runtime by setting USE_RUST_ZKP=true.

What was built

┌─────────────────────────────────────────────────────────────────┐
│  Stage 1 — Federated Averaging                                  │
│    Clients run local SGD, send deltas Δᵢ                        │
│    Server averages weighted by dataset size                     │
│    Result: 0.9904 test accuracy in 30 rounds                    │
└─────────────────────────────────────────────────────────────────┘
                            │
                            ▼  (server sees individual Δᵢ — privacy leak)
┌─────────────────────────────────────────────────────────────────┐
│  Stage 2 — Secure Aggregation + Poisoning Demonstration         │
│    Pairwise masks sᵢⱼ = −sⱼᵢ hide every individual Δᵢ           │
│    Privacy gained ✓                                             │
│    Robustness lost ✗ — gradient-reversal attack at ρ=0.3, α=3   │
│    Result: clean 0.9905 vs attack 0.098  (sustained 87–92 pp)   │
└─────────────────────────────────────────────────────────────────┘
                            │
                            ▼  (need a defence that doesn't see Δᵢ)
┌─────────────────────────────────────────────────────────────────┐
│  Stage 3 — Zero-Knowledge Norm-Bound Defence                    │
│    Each client proves ‖Δᵢ‖₂ ≤ B without revealing Δᵢ            │
│    Built from Pedersen commitments + Schnorr-OR + Fiat–Shamir   │
│    Server filters proofs that fail; aggregation continues clean │
│    Result: 0.098 → 0.9522 recovery; +27.7 % time, +0.23 % comm  │
└─────────────────────────────────────────────────────────────────┘
Stage Folder Headline output
1. FedAvg baseline 01_fedavg/ results.csv — 30-round accuracy curve, final 0.9904
2. Secure aggregation + poisoning 02_secure_aggregation/ results.csv — clean vs. attacked accuracy, 89 pp gap sustained
3. ZKP defence + overhead 03_zkp_defence/ results.csv — round-by-round time + comm; report_4_stats.json — summary
Writeups reports/ secure_aggregation_tradeoff.pdf, zkp_defence_overhead.pdf

What is actually implemented

FedAvg — 01_fedavg/

  • client.py:train_local() snapshots w_G(t), runs E epochs of local SGD with momentum 0.9 and returns the delta Δᵢ(t+1) = wᵢ(t+1) − w_G(t).
  • server.py:aggregate() computes the weighted average Δ_agg = Σᵢ (|Dᵢ| / |D|) · Δᵢ and applies the global step w_G(t+1) = w_G(t) + η · Δ_agg with η = 1.
  • main.py drives 30 rounds across 20 clients on an IID MNIST partition.

Secure aggregation + poisoning — 02_secure_aggregation/

  • secure_aggregation.py implements pairwise additive masks sᵢⱼ = −sⱼᵢ. Each client transmits its weighted update plus the sum of pairwise masks; the server sums every message and the masks telescope to zero, leaving only the aggregate learnable.
  • attack.py:AccuracyDegradationAttack flips and amplifies the local delta: Δ̃ᵢ = −α · Δᵢ.
  • main.py runs clean (ρ = 0) and attacked (ρ = 0.3, α = 3) FedAvg back-to-back on the same seed and writes both accuracy series side-by- side into one CSV.

With FedAvg's (1 − ρ) − ρ · α = −0.20 effective learning rate, the aggregate update pushes the global model backwards every round; the attacked accuracy is pinned at ≈ 0.10 (random) while the clean baseline converges past 0.99 — illustrating exactly why secure aggregation complicates robustness.

Zero-knowledge defence — 03_zkp_defence/

A non-interactive zero-knowledge range proof for the scaled integer norm n* = ⌊‖Δᵢ‖² · SCALE⌋ (SCALE = 10⁴). The construction composes four textbook primitives:

  1. Group setup. RFC 5114 1024-bit MODP safe prime p; subgroup of order q = (p − 1) / 2. g = 2 is a generator; h = g^seed mod p is a second generator with unknown DL base g.
  2. Pedersen commitment. C(x; r) = g^x · h^r mod p. Perfectly hiding, computationally binding under DLP.
  3. Schnorr-OR bit proof. A non-interactive proof that C opens to either 0 or 1, built with the Cramer–Damgård–Schoenmakers OR composition: one real Schnorr branch + one simulated branch, glued by Fiat–Shamir under SHA-256.
  4. Bit-decomposition range proof. Decompose n* into k bits and commit to each one; pick the randomness so the bit commitments combine to the norm commitment via C = Π Cᵢ^(2^i). The verifier dictates k = ⌈log₂(B² · SCALE + 1)⌉ from its policy, so the prover cannot enlarge the range.

A valid proof certifies n* ∈ [0, 2^k − 1], i.e. ‖Δᵢ‖₂ ≤ √((2^k − 1) / SCALE) ≈ B. Soundness reduces to DLP in the subgroup, zero-knowledge to Pedersen hiding plus the Schnorr-OR simulator, non-interactivity to Fiat–Shamir in the random-oracle model.

server.py:filter_updates() rejects clients whose proofs do not verify; the survivors are aggregated under the same secure-aggregation channel.

Caveat (honest-prover scope). The prover computes n* from its own Δᵢ and refuses to attach a proof when n* > 2^k − 1. Binding n* to the actual Δᵢ would require a SNARK over an arithmetic circuit for the norm computation (e.g. Groth16 on bn254 R1CS), which would change the implementation effort by an order of magnitude and is out of scope for a CPU-only laptop build target.

Optional Rust + arkworks Groth16 backend

The same ‖Δᵢ‖₂ ≤ B predicate can be discharged with a real Groth16 zk-SNARK over BN254. 03_zkp_defence/zkp_rust/ ships a PyO3 Rust extension that wraps an arkworks-based Groth16 circuit; the vendored arkworks dependency (03_zkp_defence/groth16/) is included so the build is self-contained.

# Build the Rust extension (requires Rust + maturin)
pip install maturin
cd 03_zkp_defence/zkp_rust && maturin develop

# Switch the client and server to the Rust prover/verifier at runtime
cd .. && USE_RUST_ZKP=true python main.py

zkp.py exposes ZKPProverRust / ZKPVerifierRust, which the client.py and server.py files prefer when USE_RUST_ZKP=true and fall back to the pure-Python NIZK otherwise. See 03_zkp_defence/zkp_rust/BUILD.md for the full build walk-through.


Results

Metric Value Source
FedAvg final test accuracy 0.9904 01_fedavg/results.csv
Clean / attack accuracy gap (ρ=0.3, α=3) 87–92 pp sustained over 30 rounds 02_secure_aggregation/results.csv
Accuracy recovery from ZKP filter (α=5, B=4) 0.098 → 0.9522 03_zkp_defence/report_4_stats.json
Per-round time overhead (10 rounds, 20 clients) +27.7 % (58.2 → 74.3 s/round) 03_zkp_defence/results.csv
Per-round communication overhead +0.23 % (proofs ≈ 15 KB / client) 03_zkp_defence/results.csv

Quick start

# 1. Create the conda environment (Python 3.10 + PyTorch 2.0+ + numpy)
conda env create -f environment.yml
conda activate fl-secure-aggregation-zkp

# 2. Reproduce each stage (CPU is enough; tuned defaults reproduce the numbers above)
cd 01_fedavg            && python main.py     # ~25 min on CPU
cd ../02_secure_aggregation && python main.py # ~50 min on CPU (two FedAvg runs)
cd ../03_zkp_defence    && python main.py     # ~25 min on CPU

# 3. (Optional) regenerate the PDF writeups
cd ../reports
python build_report2.py
python build_report4.py

Pass --help to any main.py for the full CLI surface (number of clients and rounds, malicious ratio, attack strength, ZKP bound B, norm type, etc.).


Project structure

.
├── README.md
├── LICENSE                     # MIT
├── .gitignore
├── environment.yml             # conda env: Python 3.10 + PyTorch 2.0+ + numpy
├── requirements.txt            # pip alternative
│
├── 01_fedavg/                  # Stage 1 — vanilla FedAvg baseline
│   ├── client.py / server.py / main.py
│   ├── model.py                #   2-conv CNN
│   ├── data_utils.py           #   MNIST loader + IID/non-IID splitter
│   └── results.csv             #   30-round accuracy curve
│
├── 02_secure_aggregation/      # Stage 2 — pairwise-mask SA + poisoning demo
│   ├── secure_aggregation.py   #   SecureAggregator (pairwise masks)
│   ├── attack.py               #   gradient-reversal model poisoning
│   ├── client.py / server.py / main.py
│   ├── model.py / data_utils.py
│   └── results.csv             #   clean vs. attacked, per round
│
├── 03_zkp_defence/             # Stage 3 — ZKP norm-bound defence + overhead
│   ├── zkp.py                  #   Pedersen + Schnorr-OR NIZK (default backend)
│   ├── server.py               #   ZKP-filtering SA server
│   ├── client.py               #   client with ZKP prover
│   ├── attack.py / secure_aggregation.py
│   ├── main.py                 #   no-ZKP vs ZKP comparison + overhead
│   ├── model.py / data_utils.py
│   ├── results.csv             #   round / time_no_zkp / time_zkp / comm_no_zkp / comm_zkp
│   ├── report_4_stats.json     #   summary stats (consumed by the report builder)
│   ├── zkp_rust/               #   optional Rust + arkworks Groth16 backend (PyO3)
│   │   ├── Cargo.toml / src/   #   maturin develop → loaded when USE_RUST_ZKP=true
│   │   └── BUILD.md            #   build instructions + troubleshooting
│   └── groth16/                #   vendored arkworks ark-groth16 dependency
│
└── reports/
    ├── build_report2.py
    ├── build_report4.py
    ├── secure_aggregation_tradeoff.pdf   # 1-page privacy↔robustness writeup
    └── zkp_defence_overhead.pdf          # 1-page ZKP design + overhead writeup

License

MIT — see LICENSE.

About

Federated learning + secure aggregation + zero-knowledge norm-bound defence

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages