Federated learning trains a shared model on data that never leaves its owners. Instead of uploading documents, chat logs or medical notes to a central cluster, each participant downloads the current model, trains on its own data for a few steps, and sends back only the resulting change in weights. A coordinating server averages those changes into the next global model and repeats.
It is tempting to read federated learning as a privacy guarantee. It is not. It is an architecture that removes one obvious risk, the central copy of raw data, while opening several new ones: updates that leak the text they were computed on, participants who poison the shared model, and a server that sees far more than the final weights. This article treats federated fine-tuning of LLMs as a security system. It covers the threat model, a LoRA-based round with code, what secure aggregation and differential privacy each buy, and poisoning.
The threat model comes first
Start with who you are defending against, because every defence in this field protects against one adversary and is useless or harmful against another.
| Adversary | What they see or control | What they want | Main defence |
|---|---|---|---|
| Curious server | every individual client update | reconstruct client text or infer attributes | secure aggregation, local or distributed DP |
| Malicious client | its own update, possibly several sybil clients | backdoor or degrade the global model | client authentication, update clipping, robust aggregation |
| Model consumer | the final model weights or API | extract or test for training examples | user-level differential privacy |
| Network attacker | traffic between clients and server | read or tamper with updates | mutual TLS, signed models |
Two of these defences pull against each other. Secure aggregation hides each update from the server, which is exactly what protects clients from a curious server. Robust aggregation needs to inspect each update to discard outliers, which is exactly what protects the model from malicious clients. You cannot have full strength of both at once, and the design choice between them should come from asking which adversary is realistic in your federation.
The threat model comes first
Start with who you are defending against, because every defence in this field protects against one adversary and is useless or harmful against another.
| Adversary | What they see or control | What they want | Main defence |
|---|---|---|---|
| Curious server | every individual client update | reconstruct client text or infer attributes | secure aggregation, local or distributed DP |
| Malicious client | its own update, possibly several sybil clients | backdoor or degrade the global model | client authentication, update clipping, robust aggregation |
| Model consumer | the final model weights or API | extract or test for training examples | user-level differential privacy |
| Network attacker | traffic between clients and server | read or tamper with updates | mutual TLS, signed models |
Two of these defences pull against each other. Secure aggregation hides each update from the server, which is exactly what protects clients from a curious server. Robust aggregation needs to inspect each update to discard outliers, which is exactly what protects the model from malicious clients. You cannot have full strength of both at once, and the design choice between them should come from asking which adversary is realistic in your federation.
One round of federated fine-tuning
Full-parameter federated training of a multi-billion-parameter model is impractical for most participants: every client would need to hold optimizer state for the whole model and ship gigabytes per round. In practice federated LLM work freezes the base model and trains a parameter-efficient adapter, most often LoRA. The base weights are distributed once; only the adapter moves each round.
One round looks like this. The coordinator samples a cohort of clients that are eligible (online, idle, on an unmetered link, or simply scheduled). Each sampled client loads the current global adapter, runs a fixed number of local optimizer steps on its own data, computes the difference between its adapter and the global one, clips that difference to a norm bound, and uploads it. The coordinator combines the deltas, optionally adds noise, applies the result to the global adapter and publishes round t+1. This is FedAvg, McMahan and co-authors' federated averaging, applied to adapter weights.
Two properties make this different from distributed data-parallel training. Clients take many local steps between synchronisations, so their models drift apart when their data differs, which it always does: one hospital is oncology, another is paediatrics. And the cohort changes every round, so the population the model is learning from is a moving sample.
Client and server code
The following is a framework-neutral sketch in PyTorch with Hugging Face PEFT. get_peft_model_state_dict returns only the adapter tensors and set_peft_model_state_dict loads them into a PEFT model, so the base model is never touched or transmitted. Transport, scheduling and authentication are left to whatever federated framework you run.
import torch
from peft import get_peft_model_state_dict, set_peft_model_state_dict
def client_update(model, global_adapter, loader, steps, lr, clip_norm):
"""Run local steps from the global adapter; return a clipped delta."""
set_peft_model_state_dict(model, global_adapter)
trainable = [p for p in model.parameters() if p.requires_grad]
opt = torch.optim.AdamW(trainable, lr=lr)
model.train()
for _, batch in zip(range(steps), loader):
loss = model(**batch).loss
loss.backward()
opt.step()
opt.zero_grad()
local = get_peft_model_state_dict(model)
delta = {k: local[k].detach().float() - global_adapter[k].float()
for k in global_adapter}
# Bound one client's influence on the round: clip the whole update in L2.
norm = torch.sqrt(sum((d ** 2).sum() for d in delta.values())).item()
scale = min(1.0, clip_norm / (norm + 1e-12))
return {k: d * scale for k, d in delta.items()}, norm
def server_round(global_adapter, deltas, clip_norm, noise_multiplier):
"""Equal-weight mean of clipped deltas, plus Gaussian noise for user-level DP."""
n = len(deltas)
mean = {k: sum(d[k] for d in deltas) / n for k in global_adapter}
if noise_multiplier > 0:
std = noise_multiplier * clip_norm / n
mean = {k: v + torch.randn_like(v) * std for k, v in mean.items()}
return {k: global_adapter[k].float() + mean[k] for k in global_adapter}Three choices in that code matter for security, not just accuracy. The clip is applied to the whole update, not per tensor, because the guarantee you want is about one client's total influence. The mean uses equal weights; weighting by local example count is common in plain FedAvg, but it lets a client that claims a huge dataset dominate the round, and it breaks the sensitivity calculation for differential privacy. And the logged pre-clip norm is your cheapest anomaly signal: a client whose norms are consistently far above the cohort is either mis-configured or attacking.
Why averaging LoRA factors is subtly wrong
LoRA represents a weight change as a product, delta_W = B @ A, where A projects down to rank r and B projects back up. Averaging the A matrices and the B matrices separately, which is what the code above does, does not give the average of the products: mean(B_i) times mean(A_i) is not mean(B_i A_i). With similar clients the error is small; with heterogeneous clients and many local steps it is a real source of instability, and adding DP noise makes it worse because noise enters both factors and is multiplied.
The simplest fix was published as FFA-LoRA (Sun and co-authors, ICLR 2024): freeze the randomly initialised A matrices, distributing them from a shared seed, and train only the zero-initialised B matrices. With A identical everywhere, mean(B_i) A equals mean(B_i A), so aggregation is exact, and each client ships half as many tensors. The cost is some expressiveness per rank, which a slightly larger rank usually recovers. If you keep vanilla LoRA, at least monitor the global model's loss on a held-out set every round rather than assuming the average is meaningful.
Worked example: bandwidth and the privacy budget
Take a Llama-2-7B-shaped model, 32 layers with hidden size 4096, and LoRA of rank 16 on the query and value projections. Each adapted 4096 by 4096 matrix gets A of 16 by 4096 and B of 4096 by 16, so 131,072 parameters. Two matrices per layer across 32 layers is 64 adapters and about 8.4 million parameters: roughly 33.5 MB per update in float32, against about 14 GB for the full model in bfloat16. With 100 clients per round the coordinator ingests about 3.4 GB per round. Freezing A halves that.
Now the privacy budget. Suppose a federation of 5,000 eligible users, a cohort of 100 per round, clip norm 1.0 and noise multiplier 1.0. The noise standard deviation on the averaged update is 1.0 times 1.0 divided by 100, or 0.01 per coordinate, which is comparable to the size of an honest averaged update on a quiet round. That is the central tension of user-level DP: noise scales with clip divided by cohort size, so small federations pay heavily. A consortium of five hospitals cannot reach a meaningful user-level guarantee where the user is a hospital; it can only protect individual patients by running DP-SGD inside each hospital, which is a different guarantee with its own accountant. Compute the actual epsilon with a privacy accountant for your sampling scheme rather than reasoning from these numbers.
What updates leak and how to stop it
A weight update is a function of the data that produced it, and research on gradient inversion has shown that training text and images can be reconstructed from individual updates, most easily when a client's batch is small and the update covers few steps. Embedding-layer updates are especially revealing for language models: if the embedding matrix is trainable, the rows that changed name the tokens that appeared. LoRA on attention projections with frozen embeddings removes that direct signal but does not make inversion impossible.
There are three layers of protection, and they stack. Secure aggregation is a cryptographic protocol (Bonawitz and co-authors described the widely used one) in which clients add pairwise masks that cancel in the sum, so the server learns only the aggregate, and dropouts are handled by secret-shared mask recovery. It protects against a curious server but does nothing about what the aggregate or the final model reveals. Differential privacy with per-client clipping and Gaussian noise, as in the code, bounds what the final model reveals about any one participant; if the noise is added by the server, you are trusting the server to add it. Trusted execution, aggregating inside an attested enclave, is an alternative or complement when clients are few and secure aggregation overhead is not worth it. None of them stops a client from memorising its own secrets into the shared model if DP is off, so membership-inference testing of the released model remains mandatory.
Poisoning and robust aggregation
Every client is a writer to the shared model. A malicious participant can train on poisoned data, or skip honest training entirely and upload a crafted delta. Bagdasaryan and co-authors showed a model-replacement attack in which one client scales its backdoor update so that it survives averaging. Targeted backdoors are the realistic threat for LLMs: the model behaves normally until a trigger phrase appears, then produces attacker-chosen output, and aggregate accuracy metrics do not move.
Defences, cheapest first:
- Norm clipping. Already in the code. It defeats naive scaling attacks and is free. It does not stop a patient attacker who stays under the bound for many rounds.
- Authenticated, rate-limited clients. Sybils are the multiplier for every attack. Tie client identity to device attestation or contract, and cap how often one identity can be sampled.
- Robust aggregation. Replace the mean with a coordinate-wise median or trimmed mean so that a minority of outliers cannot steer the result. This requires seeing individual updates, so it conflicts with secure aggregation unless performed inside an enclave.
- Canary and trigger evaluation. Before publishing each global adapter, run a fixed suite of behavioural probes and refuse to promote a model whose outputs on them change sharply.
def trimmed_mean(deltas, trim_frac=0.1):
"""Coordinate-wise trimmed mean: drop the top and bottom trim_frac per coordinate."""
out = {}
for k in deltas[0]:
stack = torch.stack([d[k] for d in deltas]) # [n_clients, ...]
n = stack.shape[0]
t = int(n * trim_frac)
ordered, _ = torch.sort(stack, dim=0)
out[k] = ordered[t:n - t].mean(dim=0)
return outThe stacked tensor costs n times the adapter size in memory, so with 100 clients and the adapter above you need over 3 GB on the aggregator, which is fine on a server and impossible inside some enclaves. Trimming also discards honest but unusual clients, which biases the model against minority data; measure per-client contribution before you turn the trim fraction up.
Operating a federation
Operating a federation is mostly about not trusting your own dashboards. Keep a held-out evaluation set at the coordinator and evaluate the global adapter every round, because client-reported training loss is unverifiable and an attacker can report anything. Version every global adapter, sign it, and make clients verify the signature before loading, so a compromised distribution channel cannot push weights. Keep the privacy accountant's running epsilon as a first-class metric with a hard stop, because the budget is spent by every round whether or not the round helped.
Plan for stragglers and dropouts. Over-sample the cohort, set a round deadline, and aggregate whatever arrived, but remember that secure aggregation must tolerate the dropout rate you actually see, and that a smaller realised cohort means more noise per coordinate if your DP calibration assumed the target size. Finally, decide up front what happens when a participant leaves the federation and asks for its contribution removed. Federated averaging has no cheap unlearning; the honest answers are retraining from a checkpoint before they joined or relying on DP to bound what remains.
Failure modes
- Calling it private without DP. Raw data stays local, but updates and the final model leak. Without clipping, noise and an accountant you have no quantitative claim.
- Trainable embeddings. Embedding-row updates reveal which tokens each client saw. Freeze them unless you truly need new vocabulary.
- Secure aggregation plus anomaly detection on individual updates. These contradict each other; one of them is not doing what you think.
- Ignoring non-IID drift. Too many local steps on skewed data makes the average worse than any client. Tune local steps per round, not just learning rate.
What to do next
Related reading on this site: DP-SGD, in depth for the per-example version of clipping and noise, data poisoning and backdoor detection for the client-side attack, membership inference attacks for testing what the final model leaks, and confidential compute for enclave-based aggregation.
- Write down the threat model: is the coordinator trusted, how many clients are there, can one party control several clients? Pick secure aggregation or robust aggregation based on that answer.
- Freeze the base model and embeddings; start with LoRA on attention projections, and consider freezing A from a shared seed so aggregation is exact.
- Clip each client update in L2 over the whole adapter and log the pre-clip norm per client per round.
- If you claim privacy, add calibrated noise, run a privacy accountant for your sampling scheme and set a hard epsilon stop.
- Build a coordinator-side evaluation suite with held-out data and behavioural trigger probes, and gate every global adapter on it before signing and publishing.
- Run a membership-inference test against the released adapter before any external use.
- Rehearse a client departure and a poisoned round: know which checkpoint you would roll back to and how long retraining takes.