Interactive note · AAAI 2025

Federated Unsupervised Domain Generalization using Global and Local Alignment of Gradients

[Paper, arXiv, Code, BibTeX]
Problem

Each client holds unlabeled data from its own domain and cannot share it. The model still has to work on a domain no client has seen.

Insight

Domain shift shows up as disagreement between gradients, and gradients are exactly what federated learning gets to see.

Fix

Clients drop batch gradients that disagree with the global direction, and the server down-weights clients that disagree. Best results on all four benchmarks.

1 · The setting

Many domains, no labels, no sharing

Picture a network of wearable activity monitors. Every device records under its own conditions, so each one sees a slightly different world, a different domain. Nobody labels their activity, and privacy rules keep the raw data on the device. Yet the model they train together should work for a new user it has never seen.

Earlier work handled these constraints one at a time: federated domain generalization assumes labels, and federated unsupervised learning ignores domain shift. This paper puts them together.

Definition 1

Federated unsupervised domain generalization is learning general representations from decentralized, unlabeled datasets, each from a different domain, when data cannot be shared.

The benchmarks make this concrete. In PACS every domain is a client except one, which is held out for testing. Pick which domain to hold out.

Leave one domain out. Three domains train as clients, each on its own device; the fourth is the unseen target. The four images are the same class, dog, from each PACS domain: what changes is the style, not the content. Quality is measured by training a linear classifier on top of the learned representation with 10% of the target's labels.
2 · Theorem 1

Different domains, different gradients

Under these rules the server never sees data, means or variances. What it does see is every client's model update, and each client sees the global update. The paper's first result says that this is enough to detect domain shift.

Model two domains as features that are correlated across domains with covariance \(\sigma\). Their mutual information, a measure of how similar the domains are, is

\[I(x_i; x_j) = -\tfrac12 \sum_{f=1}^{F} \log\!\big(1 - \sigma_{x_i^f, x_j^f}^2\big).\]
Theorem 1

For clients trained with self-supervised learning, with features modeled as standardized Gaussians and gradients approximated to first order, as the shift between two clients' domains grows, the covariance of their gradients, \(\mathrm{Cov}(g_i, g_j)\), shrinks.

Below, two domains are sampled with the similarity you choose, and the gradients of a one-layer sigmoid encoder, as in the theorem but with a simpler loss, are computed in your browser. Move the slider and watch the points line up.

0.70
Lower means more domain shift.
One feature, two domainseach dot pairs a sample from each
Gradient covariance vs mutual informationone point per run
Theorem 1, computed live. Six standardized features, a three-unit sigmoid encoder, 4,000 sample pairs per run. Left: as \(\sigma\) drops, the two domains decorrelate. Right: every \(\sigma\) you visit adds a point; less mutual information means less gradient covariance. The paper measures the same trend between real PACS domains (Section 4).

The proof runs through a first-order Taylor expansion (Lemma 2): the gradient covariance is a sum of the feature covariances \(\sigma\), weighted by products of derivatives that are always positive in this setting (Claim 1). A corollary follows: as domains drift apart, the variance of the difference of their gradients grows. So gradient alignment is a signal of domain shift that can be read without sharing raw data, and that is what FedGaLA uses. (Sharing gradients is not by itself a formal privacy guarantee.)

3 · FedGaLA

One communication round

FedGaLA aligns gradients in two places.

On each client, training is self-supervised (SimCLR). For every batch and every layer, the gradient is compared with a reference: the global model's change over the last round, \(\hat g_{est} = \theta^{(t)} - \theta^{(t-1)}\). If its cosine with the reference is below a threshold \(\tau\), the batch gradient is discarded. Scaling it down would not help, since cosine ignores scale.

On the server, each client's update gets a weight \(w_i = \tfrac12\big(\cos(\hat g_i, \hat g) + 1\big)\), normalized across clients, where \(\hat g\) is the current aggregate. The aggregate is recomputed with the new weights, three times. Dropping a whole client would throw away a domain, so the server weighs softly.

Here is one round. Two domains are typical; make the third one different and press Play round.

1.2
0.00
The paper finds \(\tau = 0\) works best.
On the clients: local alignment32 batch gradients each
On the server: global alignmentclient updates and the aggregate
Batch gradients discardedper client
Aggregation weightsplain average vs FedGaLA
One round, in two dimensions. Up is the direction the domains share; the dashed arrow is last round's global update, the reference. Left: each client's batch gradients, scattered by noise around its own domain's direction. Gradients on the wrong side of the threshold fade out. Right: what each client sends, the plain average FedAvg would take (grey), and the FedGaLA aggregate (purple). An unusual domain loses most of its disagreeing gradients locally, then gets a slightly smaller say at the server, so it cannot drag the shared model toward its quirks. A toy for intuition; the real model has millions of parameters and aligns each layer separately.

The global weights are gentle by design: \(w_i\) only falls to zero when a client points exactly against the aggregate. Most of the correction happens on the clients, where unaligned batches are simply skipped. In the paper's ablation the two parts are complementary: either one alone barely changes PACS accuracy, while together they add 1.4 points.

4 · Experiments

What the experiments show

Four benchmarks

Best on PACS, DomainNet, Office-Home and TerraInc

ResNet-18 trained from scratch, 100 rounds. Linear evaluation with 10% of the target's labels, averaged over held-out domains. Every baseline is a self-supervised method made federated with FedAvg; FedGaLA is FedSimCLR plus the two alignments.

Method PACS DomainNet Office-Home TerraInc
FedEMA 41.9 32.4 13.5 53.2
FedBYOL 44.2 31.8 13.8 54.3
FedMoCo 42.1 27.2 10.7 45.7
FedSimSiam 39.8 36.9 18.9 47.9
FedSimCLR 58.6 39.5 22.0 55.1
FedGaLA 60.0 41.1 23.0 56.7
The theory, measured

Pairs of PACS domains

Gradient covariance against mutual information for pairs of PACS domains

Gradient covariance against mutual information for each pair of domains. Apart from one outlier, more shift means less covariance, as Theorem 1 predicts.

Clients converge

Fewer gradients discarded over time

Share of discarded local gradients falling over 100 communication rounds

The share of discarded batch gradients falls from about 68% to 37% over 100 rounds: the clients learn features that agree with the global model.

Ablation, PACS

Both alignments matter

Variant Accuracy
FedGaLA 60.0
without global alignment 58.8
without local alignment 58.5
without both (FedSimCLR) 58.6

Neither alignment helps much alone; together they add 1.4 points.

Federated beats centralized

Better than training on pooled data

PACS, 10% labels. The centralized methods see all domains in one place with no privacy constraint; FedGaLA still leads by a wide margin, in line with prior evidence that federation helps domain generalization.

Cite

BibTeX

@inproceedings{pourpanah2025fedgala,
  title     = {Federated Unsupervised Domain Generalization
               using Global and Local Alignment of Gradients},
  author    = {Pourpanah, Farhad and Molahasani, Mahdiyar and
               Soltany, Milad and Greenspan, Michael and
               Etemad, Ali},
  booktitle = {Proceedings of the AAAI Conference on
               Artificial Intelligence},
  pages     = {19948--19958},
  year      = {2025}
}

* Equal contribution. The models in §2 and §3 are small toys built and run in your browser. The numbers and plots in §4 are from the paper.

← All writing