FedTGP: Trainable Global Prototypes with Adaptive-Margin-Enhanced Contrastive Learning for Data and Model Heterogeneity in Federated Learning

Jianqing ZhangYang LiuYang HuaJian Cao

article2024AAAI134 citations

Proposes a heterogeneous federated learning framework that replaces naive prototype averaging with server-side trainable global prototypes optimized through adaptive-margin contrastive learning, boosting classification accuracy across diverse client models while preserving communication efficiency and privacy.

Listen

Modern distributed artificial intelligence systems increasingly face two simultaneous challenges: protecting the privacy and intellectual property of participating clients while accommodating diverse local model architectures and skewed data distributions. While heterogeneous federated learning allows participants to collaborate without exposing their private models or raw data, existing lightweight prototype-based methods rely on naive weighted-averaging to aggregate class summaries. This conventional averaging causes separation margins between classes to shrink, degrades performance, and inadvertently leaks sensitive local data distribution information to the central server.

The article introduces and evaluates FedTGP, a novel framework designed to overcome these limitations by learning trainable global prototypes on the central server through adaptive-margin-enhanced contrastive learning. The primary objective is to demonstrate that FedTGP can enhance class separability and overall model accuracy across heterogeneous clients without increasing communication overhead or exposing private model weights and data distributions.

To establish credibility, the authors conducted extensive empirical evaluations across four standard image classification datasets (CIFAR-10, CIFAR-100, Flowers102, and Tiny-ImageNet) involving up to 100 simulated clients. The testing environment spanned twelve heterogeneous model architectures, including standard convolutional neural networks, GoogleNet, MobileNet v2, and various ResNet variants. The evaluations tested both pathological and practical non-uniform data splits, comparing FedTGP against six state-of-the-art baselines across varying feature dimensions, local training epochs, and partial client participation levels.

The empirical findings demonstrate substantial performance gains. First, FedTGP outperformed all baseline methods across every evaluated dataset, exceeding its primary prototype-based predecessor, FedProto, by up to 13.85% in standard setups and outperforming other leading approaches by up to 9.08% in accuracy. Second, in extreme heterogeneity scenarios where both feature extractors and classifiers differed across clients, FedTGP maintained its robustness, outperforming FedProto by 18.96%. Third, the framework demonstrated resilience to network scaling; under partial participation with 100 clients, FedTGP maintained a 5.48% accuracy lead over the best alternative. Finally, FedTGP maintained a minimal communication footprint of approximately 1.48 million parameters per round—substantially lower than distillation-based methods requiring up to 36.99 million parameters—while eliminating the need for clients to share sample counts.

These results demonstrate that server-side trainable prototype learning effectively decouples local model optimization from global aggregation while strengthening inter-class decision boundaries. For organizations deploying federated systems, adopting this approach mitigates data leakage risks, preserves proprietary model intellectual property, and reduces network bandwidth costs without sacrificing model accuracy, especially in complex classification tasks with many categories.

Organizations implementing federated learning across heterogeneous devices should consider adopting adaptive-margin prototype frameworks over standard averaging protocols or bandwidth-heavy knowledge distillation methods. When deploying FedTGP, engineering teams should tune server training iterations and distance thresholds to balance computational efficiency with prototype stability, noting that modest server training budgets (e.g., 100 epochs) yield near-optimal accuracy. While the experimental evidence is highly consistent across evaluated computer vision benchmarks, stakeholders should exercise cautious optimism and conduct targeted pilot tests before deploying the framework to non-vision modalities, real-world network latency conditions, or settings with extreme label corruption.

  • Paper: Model-Contrastive Federated Learning, Qinbin Li et al. (2021). Introduces model-contrastive learning in federated settings to correct client drift under data heterogeneity, establishing the foundational paradigm that FedTGP extends to server-side trainable prototypes.
  • Paper: Communication-Efficient Learning of Deep Networks from Decentralized Data, H. B. McMahan et al. (2016). Introduces FederatedAveraging (FedAvg), the canonical baseline for distributed local optimization whose limitations under severe data and model heterogeneity motivated prototype-based federated learning.
  • Paper: Federated Optimization in Heterogeneous Networks, Tian Li et al. (2018). Presents FedProx to handle statistical and system heterogeneity via proximal regularization, framing the core challenges of client drift that prototype-based methods address.
  • Paper: Ensemble Distillation for Robust Model Fusion in Federated Learning, Tao Lin et al. (2020). Develops ensemble distillation for model-heterogeneous federated learning, providing essential background on training across disparate client architectures without parameter averaging.
  • Paper: Federated Learning with Personalization Layers, Manoj Ghuhan Arivazhagan et al. (2019). Proposes decoupling representation feature extractors from classifier heads in heterogeneous federated learning, laying groundwork for prototype sharing across non-identical clients.
  • Paper: Federated Learning with Non-IID Data, Yue Zhao et al. (2018). Analyzes the mathematical mechanisms of weight divergence caused by non-IID client partitions, motivating the shift toward representation and prototype aggregation.

No sufficiently relevant recommendations were found.

Cover for FedTGP: Trainable Global Prototypes with Adaptive-Margin-Enhanced Contrastive Learning for Data and Model Heterogeneity in Federated Learning

Abstract

Recently, Heterogeneous Federated Learning (HtFL) has attracted attention due to its ability to support heterogeneous models and data. To reduce the high communication cost of transmitting model parameters, a major challenge in HtFL, prototype-based HtFL methods are proposed to solely share class representatives, a.k.a, prototypes, among heterogeneous clients while maintaining the privacy of clients' models. However, these prototypes are naively aggregated into global prototypes on the server using weighted averaging, resulting in suboptimal global knowledge which negatively impacts the performance of clients. To overcome this challenge, we introduce a novel HtFL approach called FedTGP, which leverages our Adaptive-margin-enhanced Contrastive Learning (ACL) to learn Trainable Global Prototypes (TGP) on the server. By incorporating ACL, our approach enhances prototype separability while preserving semantic meaning. Extensive experiments with twelve heterogeneous models demonstrate that our FedTGP surpasses state-of-the-art methods by up to 9.08% in accuracy while maintaining the communication and privacy advantages of prototype-based HtFL. Our code is available at https://github.com/TsingZ0/FedTGP.

Table of Contents

  • Introduction
  • Related Work
  • Heterogeneous Federated Learning
  • Trainable Prototype Learning
  • Method Problem Statement and Motivation
  • Trainable Global Prototypes
  • Adaptive-Margin-Enhanced Contrastive Learning
  • FedTGP Framework
  • Experiments
  • Setup
  • Performance
  • Impact of Model Heterogeneity
  • Partial Participation with More Clients
  • Impact of Number of Client Training Epochs
  • Impact of Feature Dimensions
  • Communication Cost
  • Ablation and Hyperparameter Study
  • Conclusion
  • Acknowledgments
  • References

Knowls

  1. Knowl 1 — FedTGP Collaborative Training Algorithm

    algorithm

    The FedTGP algorithm trains heterogeneous client models collaboratively without sharing model parameters, auxiliary models, or private sample count distributions. Instead, the server maintains trainable global prototypes parameterized by class vectors and a shared feed-forward network, optimizing them with an adaptive-margin-enhanced contrastive objective over received client prototypes.

    Input: Number of clients MM, private client datasets {Di}i=1M\{D_i\}_{i=1}^M, trainable global prototype parameters P^={Pˊc}c=1C\hat{\mathcal{P}} = \{\acute{P}^c\}_{c=1}^C and network parameters θF\theta_F, client learning rate η\eta, communication rounds TT, server training epochs SS, margin cap τ\tau, regularization weight λ\lambda, client sampling ratio ρ\rho.
    Output: Trained client model parameters {(θi,wi)}i=1M\{(\theta_i, w_i)\}_{i=1}^M.
    for iteration t=1,…,Tt = 1, \dots, T do
        Server randomly samples a subset of clients ItI^t with size ρM\rho M.
        Server computes current global prototypes P^c=F(Pˊc;θF)\hat{P}^c = F(\acute{P}^c; \theta_F) for all c∈[C]c \in [C].
        Server broadcasts P^={P^c}c=1C\hat{\mathcal{P}} = \{\hat{P}^c\}_{c=1}^C to participating clients ItI^t.
        for each client i∈Iti \in I^t in parallel do
            Update local feature extractor θi\theta_i and classifier wiw_i by gradient descent on the local objective:
            Li=E(x,y)∼Diℓ(hi(fi(x;θi);wi),y)+λEc∼Ci∥Pic−P^c∥2\mathcal{L}_i = \mathbb{E}_{(x, y) \sim D_i} \ell(h_i(f_i(x; \theta_i); w_i), y) + \lambda \mathbb{E}_{c \sim \mathcal{C}_i} \|P_i^c - \hat{P}^c\|_2
            Compute local class prototypes for all classes c∈Cic \in \mathcal{C}_i present on client ii:
            Pic=E(x,c)∼Di,cfi(x;θi)P_i^c = \mathbb{E}_{(x, c) \sim D_{i, c}} f_i(x; \theta_i)
            Send local prototypes Pi={Pic}c∈Ci\mathcal{P}_i = \{P_i^c\}_{c \in \mathcal{C}_i} to the server.
        Server calculates unweighted cluster centers for each class c∈[C]c \in [C]:
        Qtc=1∣Ptc∣∑i∈ItPicQ_t^c = \frac{1}{|\mathcal{P}_t^c|} \sum_{i \in I^t} P_i^c
        Server determines the adaptive margin:
        δ(t)=min⁡(max⁡c≠c′∥Qtc−Qtc′∥2,τ)\delta(t) = \min\left( \max_{c \neq c'} \|Q_t^c - Q_t^{c'}\|_2, \tau \right)
        for server epoch s=1,…,Ss = 1, \dots, S do
            Server updates {Pˊc}c=1C\{\acute{P}^c\}_{c=1}^C and θF\theta_F by minimizing the adaptive-margin contrastive loss ∑c=1CLPc\sum_{c=1}^C \mathcal{L}_P^c.
    return Client models {(θi,wi)}i=1M\{(\theta_i, w_i)\}_{i=1}^M.
  2. Knowl 2 — Trainable Global Prototypes Representation and Architecture

    model/method

    In FedTGP, global prototypes are not computed via fixed statistical aggregation of client prototypes. Instead, the server maintains a set of Trainable Global Prototypes (TGP) P^={P^c}c=1C\hat{\mathcal{P}} = \{\hat{P}^c\}_{c=1}^C, where CC is the total number of classes.

    Each class c∈[C]c \in [C] is assigned a trainable latent vector Pˊc∈RK\acute{P}^c \in \mathbb{R}^K, where KK is the feature dimension. To enhance representational capacity, a shared neural network F(⋅;θF)F(\cdot; \theta_F) parameterized by θF\theta_F processes these vectors to produce the final global class prototypes:

    P^c=F(Pˊc;θF)∈RK,∀c∈[C]\hat{P}^c = F(\acute{P}^c; \theta_F) \in \mathbb{R}^K, \quad \forall c \in [C]

    The network FF consists of two Fully-Connected (FC) layers with an intermediate ReLU activation function (K→K→KK \to K \to K). The parameters {Pˊc}c=1C\{\acute{P}^c\}_{c=1}^C and θF\theta_F are situated entirely on the server and optimized jointly during server-side training without accessing client raw data or client model weights.

  3. Knowl 3 — Adaptive-Margin-Enhanced Contrastive Loss for Global Prototypes

    equation

    To guarantee that trainable global prototypes P^={P^c}c=1C\hat{\mathcal{P}} = \{\hat{P}^c\}_{c=1}^C maintain class semantics while enforcing inter-class separation, the server minimizes an Adaptive-margin-enhanced Contrastive Learning (ACL) objective over the participating clients ItI^t at iteration tt:

    min⁡P^∑c=1CLPc\min_{\hat{\mathcal{P}}} \sum_{c=1}^C \mathcal{L}_P^c

    LPc=∑i∈It−log⁡exp⁡(−(ϕ(Pic,P^c)+δ(t)))exp⁡(−(ϕ(Pic,P^c)+δ(t)))+∑c′≠cexp⁡(−ϕ(Pic,P^c′))\mathcal{L}_P^c = \sum_{i \in I^t} -\log \frac{\exp\left(-\left(\phi(P_i^c, \hat{P}^c) + \delta(t)\right)\right)}{\exp\left(-\left(\phi(P_i^c, \hat{P}^c) + \delta(t)\right)\right) + \sum_{c' \neq c} \exp\left(-\phi(P_i^c, \hat{P}^{c'})\right)}

    where ϕ(u,v)=∥u−v∥2\phi(u, v) = \|u - v\|_2 denotes the Euclidean distance, PicP_i^c is the prototype for class cc from client ii, and c′∈[C]∖{c}c' \in [C] \setminus \{c\}.

    The margin δ(t)\delta(t) is dynamically adjusted at each communication round tt to match the maximum cluster margin among client prototypes across different classes, capped by a threshold hyperparameter τ>0\tau > 0:

    δ(t)=min⁡(max⁡c∈[C],c′∈[C],c≠c′ϕ(Qtc,Qtc′),τ)\delta(t) = \min\left( \max_{c \in [C], c' \in [C], c \neq c'} \phi(Q_t^c, Q_t^{c'}), \tau \right)

    where Qtc=1∣Ptc∣∑i∈ItPicQ_t^c = \frac{1}{|\mathcal{P}_t^c|} \sum_{i \in I^t} P_i^c is the unweighted centroid of the received prototypes for class cc across active clients Ptc={Pic}i∈It\mathcal{P}_t^c = \{P_i^c\}_{i \in I^t}.

  4. Knowl 4 — Client Local Optimization and Prototype-Based Inference

    model/method

    In FedTGP, each client i∈{1,…,M}i \in \{1, \dots, M\} possesses a private dataset DiD_i and a local model decomposed into a feature extractor fi(⋅;θi):RD→RKf_i(\cdot; \theta_i): \mathbb{R}^D \to \mathbb{R}^K and a classifier header hi(⋅;wi):RK→RCh_i(\cdot; w_i): \mathbb{R}^K \to \mathbb{R}^C.

    Each client computes its local prototype for each class c∈Cic \in \mathcal{C}_i present in DiD_i by averaging the extracted features over the class subset Di,cD_{i,c}:

    Pic=E(x,c)∼Di,cfi(x;θi)P_i^c = \mathbb{E}_{(x, c) \sim D_{i, c}} f_i(x; \theta_i)

    Local model parameters θi\theta_i and wiw_i are trained by minimizing the joint loss Li\mathcal{L}_i:

    Li=E(x,y)∼Diℓ(hi(fi(x;θi);wi),y)+λEc∼Ciϕ(Pic,P^c)\mathcal{L}_i = \mathbb{E}_{(x, y) \sim D_i} \ell(h_i(f_i(x; \theta_i); w_i), y) + \lambda \mathbb{E}_{c \sim \mathcal{C}_i} \phi(P_i^c, \hat{P}^c)

    where ℓ\ell is the task supervised loss (such as cross-entropy), ϕ(u,v)=∥u−v∥2\phi(u, v) = \|u - v\|_2 is the Euclidean distance, P^c\hat{P}^c is the trainable global prototype received from the server, and λ\lambda is a regularization hyperparameter (set to λ=0.1\lambda = 0.1).

    For inference on client ii, an input sample xx is mapped to fi(x;θi)f_i(x; \theta_i) and assigned to the class corresponding to the nearest global prototype in Euclidean space:

    y^=arg⁡min⁡c∈[C]∥fi(x;θi)−P^c∥2\hat{y} = \arg\min_{c \in [C]} \| f_i(x; \theta_i) - \hat{P}^c \|_2

  5. Knowl 5 — Prototype Margin Shrink Phenomenon

    definition

    In prototype-based heterogeneous federated learning schemes (e.g., FedProto), the global prototype Pˉc\bar{P}^c for class cc is computed by weighted averaging over client prototypes PicP_i^c using private data counts ∣Di,c∣|D_{i, c}|/NcN_c as weights:

    Pˉc=1∣Nc∣∑i∈Nc∣Di,c∣NcPic\bar{P}^c = \frac{1}{|\mathcal{N}_c|} \sum_{i \in \mathcal{N}_c} \frac{|D_{i,c}|}{N_c} P_i^c

    Prototype Margin Shrink refers to the degradation phenomenon where weighted averaging across models with heterogeneous architectures and varying feature quality produces global prototypes with inter-class separation margins min⁡c′≠c∥Pˉc−Pˉc′∥2\min_{c' \neq c} \|\bar{P}^c - \bar{P}^{c'}\|_2 that are substantially smaller than the maximum separation margins achieved by individual well-performing clients. This occurs because clients with inferior feature extraction capability may hold large data quantities and dominate the weighted average, collapsing inter-class boundaries and degrading client model training.

  6. Knowl 6 — Classification Accuracy Benchmark Across Heterogeneous Data Settings

    data/table

    The classification performance of FedTGP was evaluated against six Heterogeneous Federated Learning (HtFL) baselines using an ensemble of eight heterogeneous feature extractor architectures (HtFE8: 4-layer CNN, GoogleNet, MobileNet_v2, ResNet18, ResNet34, ResNet50, ResNet101, ResNet152) across four image datasets. Two non-IID data distributions were tested: pathological (non-redundant class allocation: 2 classes/client on Cifar10, 10 on Cifar100, 10 on Flowers102, 20 on Tiny-ImageNet) and practical (Dirichlet distribution Dir(β)\text{Dir}(\beta) with β=0.1\beta = 0.1).

    Settings Pathological Setting Practical Setting
    Datasets Cifar10 Cifar100 Flowers102 Tiny-ImageNet Cifar10 Cifar100 Flowers102 Tiny-ImageNet
    LG-FedAvg 86.82±\pm0.26 57.01±\pm0.66 58.88±\pm0.28 32.04±\pm0.17 84.55±\pm0.51 40.65±\pm0.07 45.93±\pm0.48 24.06±\pm0.10
    FedGen 82.83±\pm0.65 58.26±\pm0.36 59.90±\pm0.15 29.80±\pm1.11 82.55±\pm0.49 38.73±\pm0.14 45.30±\pm0.17 19.60±\pm0.08
    FML 87.06±\pm0.24 55.15±\pm0.14 57.79±\pm0.31 31.38±\pm0.15 85.88±\pm0.08 39.86±\pm0.25 46.08±\pm0.53 24.25±\pm0.14
    FedKD 87.32±\pm0.31 56.56±\pm0.27 54.82±\pm0.35 32.64±\pm0.36 86.45±\pm0.10 40.56±\pm0.31 48.52±\pm0.28 25.51±\pm0.35
    FedDistill 87.24±\pm0.06 56.99±\pm0.27 58.51±\pm0.34 31.49±\pm0.38 86.01±\pm0.31 41.54±\pm0.08 49.13±\pm0.85 24.87±\pm0.31
    FedProto 83.39±\pm0.15 53.59±\pm0.29 55.13±\pm0.17 29.28±\pm0.36 82.07±\pm1.64 36.34±\pm0.28 41.21±\pm0.22 19.01±\pm0.10
    FedTGP 90.02±\pm0.30 61.86±\pm0.30 68.98±\pm0.43 34.56±\pm0.27 88.15±\pm0.43 46.94±\pm0.12 53.68±\pm0.31 27.37±\pm0.12

    FedTGP outperforms all baseline methods across all eight benchmark configurations, exceeding FedProto by up to 13.85% (Flowers102 pathological) and achieving up to 9.08% improvement over the best non-prototype baseline (Flowers102 pathological).

  7. Knowl 7 — Robustness to Feature Extractor and Classifier Heterogeneity and Client Scaling

    data/table

    Test accuracy (%) on Cifar100 under practical non-IID distribution (Dir(0.1)\text{Dir}(0.1)) was evaluated across varying degrees of model heterogeneity, heterogeneous classifier setups, and scaled client cohorts with partial participation (ρ=0.5\rho = 0.5):

    • Heterogeneous Feature Extractor Groups: HtFE2 (4-layer CNN, ResNet18), HtFE3 (ResNet10, ResNet18, ResNet34), HtFE4 (4-layer CNN, GoogleNet, MobileNet_v2, ResNet18), HtFE9 (ResNet4, ResNet6, ResNet8, ResNet10, ResNet18, ResNet34, ResNet50, ResNet101, ResNet152).
    • Classifier Heterogeneity: Res34-HtC4 (homogeneous ResNet34 extractors with 4 heterogeneous classifier heads), HtFE8-HtC4 (8 heterogeneous feature extractors and 4 heterogeneous classifiers).
    • Client Scaling: 50 and 100 clients under HtFE8 with client sampling ratio ρ=0.5\rho = 0.5.
    Settings Heterogeneous Feature Extractors Heterogeneous Classifiers Large Client Amount
    HtFE2 HtFE3 HtFE4 HtFE9 Res34-HtC4 HtFE8-HtC4 50 Clients 100 Clients
    LG-FedAvg 46.61±\pm0.24 45.56±\pm0.37 43.91±\pm0.16 42.04±\pm0.26 — — 37.81±\pm0.12 35.14±\pm0.47
    FedGen 43.92±\pm0.11 43.65±\pm0.43 40.47±\pm1.09 40.28±\pm0.54 — — 37.95±\pm0.25 34.52±\pm0.31
    FML 45.94±\pm0.16 43.05±\pm0.06 43.00±\pm0.08 42.41±\pm0.28 41.03±\pm0.20 39.23±\pm0.42 38.47±\pm0.14 36.09±\pm0.28
    FedKD 46.33±\pm0.24 43.16±\pm0.49 43.21±\pm0.37 42.15±\pm0.36 39.77±\pm0.42 40.59±\pm0.51 38.25±\pm0.41 35.62±\pm0.55
    FedDistill 46.88±\pm0.13 43.53±\pm0.21 43.56±\pm0.14 42.09±\pm0.20 44.72±\pm0.13 41.67±\pm0.06 38.51±\pm0.36 36.06±\pm0.24
    FedProto 43.97±\pm0.18 38.14±\pm0.64 34.67±\pm0.55 32.74±\pm0.82 32.26±\pm0.18 25.57±\pm0.72 33.03±\pm0.42 28.95±\pm0.51
    FedTGP 49.82±\pm0.29 49.65±\pm0.37 46.54±\pm0.14 48.05±\pm0.19 48.18±\pm0.27 44.53±\pm0.16 43.17±\pm0.23 41.57±\pm0.30

    When increasing model heterogeneity from HtFE2 to HtFE9, FedTGP accuracy drops by only 1.77% (from 49.82% to 48.05%), whereas baseline performances drop by 3.53% to 15.04% (FedProto drops from 43.97% to 32.74%). In the fully heterogeneous scenario HtFE8-HtC4, FedTGP outperforms FedProto by 18.96%.

  8. Knowl 8 — Per-Iteration Communication Overhead Comparison

    data/table

    Communication costs per iteration for MM clients under HtFE8 on Cifar100 in the practical setting (K=512K=512, C=100C=100) are determined by whether model parameters, auxiliary models, logits, or prototypes are exchanged.

    Method Theoretical Overhead Practice Overhead
    LG-FedAvg ∑i=1M∣wi∣×2\sum_{i=1}^M |w_i| \times 2 2.05 MB
    FedGen ∑i=1M(∣wi∣×2+∣Θ∣)\sum_{i=1}^M (|w_i| \times 2 + |\Theta|) 8.69 MB
    FML M×(∣θg∣+∣wg∣)×2M \times (|\theta_g| + |w_g|) \times 2 36.99 MB
    FedKD M×(∣θg∣+∣wg∣)×2×rM \times (|\theta_g| + |w_g|) \times 2 \times r 33.04 MB
    FedDistill ∑i=1MC×(Ci+C)\sum_{i=1}^M C \times (\mathcal{C}_i + C) 0.29 MB
    FedProto ∑i=1MK×(Ci+C)\sum_{i=1}^M K \times (\mathcal{C}_i + C) 1.48 MB
    FedTGP ∑i=1MK×(Ci+C)\sum_{i=1}^M K \times (\mathcal{C}_i + C) 1.48 MB

    Here, Θ\Theta denotes parameters of the auxiliary generator in FedGen; θg\theta_g and wgw_g are parameters of the auxiliary feature extractor and classifier in FML and FedKD; rr is the SVD compression rate in FedKD (∣θg∣≫K×C|\theta_g| \gg K \times C); Ci\mathcal{C}_i is the number of local classes on client ii. FedTGP incurs the exact same lightweight communication overhead as FedProto (1.48 MB1.48\text{ MB}), which is 25×25\times smaller than FML and 22.3×22.3\times smaller than FedKD.

  9. Knowl 9 — Component Ablations on Server Contrastive Learning and Model Architecture

    data/table

    An ablation study on Cifar100, Flowers102, and Tiny-ImageNet under the practical non-IID setting with HtFE8 evaluates the contributions of individual FedTGP components:

    • SCL: Standard Contrastive Loss without margin (replacing ACL with standard contrastive loss).
    • FM: Fixed Margin contrastive loss (using constant δ\delta rather than adaptive δ(t)\delta(t)).
    • w/o F: Trainable Global Prototypes without the feed-forward processing model F(⋅;θF)F(\cdot; \theta_F) (directly optimizing latent vectors {Pˊc}c=1C\{\acute{P}^c\}_{c=1}^C).
    • Proto: Baseline FedProto (weighted average of client prototypes).
    • TGP: Full FedTGP system.
    Dataset SCL FM w/o F FedProto FedTGP
    Cifar100 40.11% 43.46% 40.37% 36.34% 46.94%
    Flowers102 46.81% 52.03% 49.39% 41.21% 53.68%
    Tiny-ImageNet 22.26% 26.13% 23.12% 19.01% 27.37%

    On Cifar100, introducing contrastive optimization (SCL) increases accuracy over FedProto by 3.77%, adding a fixed margin (FM) improves it by 7.12%, while the full adaptive margin mechanism in FedTGP yields a 10.60% improvement over FedProto. Omitting the processing network FF reduces accuracy by 6.57% on Cifar100, showing that both the non-linear transformation and adaptive margin are critical.

  10. Knowl 10 — Sensitivity to Margin Threshold Cap and Server Training Epochs

    data/table

    The influence of the maximum margin cap τ\tau and the number of server optimization epochs SS on test accuracy (%) was examined on Cifar100 in the practical setting under the HtFE8 model group (default parameters τ=100,S=100\tau = 100, S = 100):

    Different τ\tau Different SS
    1 10 100 1000 1 10 100 1000
    43.23% 44.81% 46.94% 46.09% 43.41% 44.62% 46.94% 47.01%

    Accuracy increases as τ\tau increases from 11 to 100100, but drops slightly at τ=1000\tau = 1000 due to excessive margin values causing instability in prototype guidance on clients. Accuracy increases monotonically with server training epochs SS; however, because the gain from S=100S=100 to S=1000S=1000 is only 0.07%, S=100S=100 achieves the best computation-performance trade-off.

Coverage note — None was omitted; all primary methodological components, theoretical definitions, formal equations, algorithmic steps, benchmark tables, and hyperparameter/ablation analyses are covered.

References

  1. 1.Chen, H.-Y.; and Chao, W.-L. 2021. On Bridging Generic and Personalized Federated Learning for Image Classification. In ICLR.
  2. 2.Choi, H.; Som, A.; and Turaga, P. 2020. AMC-loss: Angular margin contrastive loss for improved explainability in image classification. In CVPR Workshop.
  3. 3.Chrabaszcz, P.; Loshchilov, I.; and Hutter, F. 2017. A Downsampled Variant of Imagenet as an Alternative to the Cifar Datasets. arXiv preprint arXiv:1707.08819.
  4. 4.Collins, L.; Hassani, H.; Mokhtari, A.; and Shakkottai, S. 2021. Exploiting Shared Representations for Personalized Federated Learning. In ICML.
  5. 5.Deng, J.; Guo, J.; Xue, N.; and Zafeiriou, S. 2019. Arcface: Additive angular margin loss for deep face recognition. In CVPR.
  6. 6.Diao, E.; Ding, J.; and Tarokh, V. 2020. HeteroFL: Computation and Communication Efficient Federated Learning for Heterogeneous Clients. In ICLR.
  7. 7.Hayat, M.; Khan, S.; Zamir, S. W.; Shen, J.; and Shao, L. 2019. Gaussian affinity for max-margin class imbalanced learning. In ICCV.
  8. 8.He, K.; Zhang, X.; Ren, S.; and Sun, J. 2016. Deep Residual Learning for Image Recognition. In CVPR.
  9. 9.Hinton, G.; Vinyals, O.; and Dean, J. 2015. Distilling the knowledge in a neural network. arXiv preprint arXiv:1503.02531.
  10. 10.Horvath, S.; Laskaridis, S.; Almeida, M.; Leontiadis, I.; Venieris, S.; and Lane, N. 2021. Fjord: Fair and accurate federated learning under heterogeneous targets with ordered dropout. NeurIPS.
  11. 11.Jeong, E.; Oh, S.; Kim, H.; Park, J.; Bennis, M.; and Kim, S.L. 2018. Communication-efficient on-device machine learning: Federated distillation and augmentation under non-iid private data. arXiv preprint arXiv:1811.11479.
  12. 12.Jin, X.-B.; Liu, C.-L.; and Hou, X. 2010. Regularized margin-based conditional log-likelihood loss for prototype learning. Pattern Recognition, 43(7): 2428–2438.
  13. 13.Kairouz, P.; McMahan, H. B.; Avent, B.; Bellet, A.; Bennis, M.; Bhagoji, A. N.; Bonawitz, K.; Charles, Z.; Cormode, G.; Cummings, R.; et al. 2019. Advances and Open Problems in Federated Learning. arXiv preprint arXiv:1912.04977.
  14. 14.Kim, T.; and Kim, C. 2020. Attract, perturb, and explore: Learning a feature alignment network for semi-supervised domain adaptation. In ECCV.
  15. 15.Krizhevsky, A.; and Geoffrey, H. 2009. Learning Multiple Layers of Features From Tiny Images. Technical Report.
  16. 16.Li, D.; and Wang, J. 2019. Fedmd: Heterogenous federated learning via model distillation. arXiv preprint arXiv:1910.03581.
  17. 17.Li, Q.; Diao, Y.; Chen, Q.; and He, B. 2022. Federated Learning on Non-IID Data Silos: An Experimental Study. In ICDE.
  18. 18.Li, Q.; He, B.; and Song, D. 2021. Model-Contrastive Federated Learning. In CVPR.
  19. 19.Li, Q.; Wen, Z.; Wu, Z.; Hu, S.; Wang, N.; Li, Y.; Liu, X.; and He, B. 2021a. A Survey on Federated Learning Systems: Vision, Hype and Reality for Data Privacy and Protection. IEEE Transactions on Knowledge and Data Engineering.
  20. 20.Li, T.; Hu, S.; Beirami, A.; and Smith, V. 2021b. Ditto: Fair and Robust Federated Learning Through Personalization. In ICML.
  21. 21.Li, T.; Sahu, A. K.; Talwalkar, A.; and Smith, V. 2020. Federated Learning: Challenges, Methods, and Future Directions. IEEE Signal Processing Magazine, 37(3): 50–60.
  22. 22.Li, Z.; Shang, X.; He, R.; Lin, T.; and Wu, C. 2023a. No Fear of Classifier Biases: Neural Collapse Inspired Federated Learning with Synthetic and Fixed Classifier. arXiv preprint arXiv:2303.10058.
  23. 23.Li, Z.; Wang, X.; Robertson, N. M.; Clifton, D. A.; Meinel, C.; and Yang, H. 2023b. SMKD: Selective Mutual Knowledge Distillation. In IJCNN.
  24. 24.Liang, P. P.; Liu, T.; Ziyin, L.; Allen, N. B.; Auerbach, R. P.; Brent, D.; Salakhutdinov, R.; and Morency, L.-P. 2020. Think locally, act globally: Federated learning with local and global representations. arXiv preprint arXiv:2001.01523.
  25. 25.Liao, Y.; Ma, L.; Zhou, B.; Zhao, X.; and Xie, F. 2023. DraftFed: A Draft-Based Personalized Federated Learning Approach for Heterogeneous Convolutional Neural Networks. IEEE Transactions on Mobile Computing.
  26. 26.Lin, T.; Kong, L.; Stich, S. U.; and Jaggi, M. 2020. Ensemble distillation for robust model fusion in federated learning. NeurIPS.
  27. 27.Luo, M.; Chen, F.; Hu, D.; Zhang, Y.; Liang, J.; and Feng, J. 2021. No Fear of Heterogeneity: Classifier Calibration for Federated Learning with Non-IID data. In NeurIPS.
  28. 28.Ma, X.; Zhang, J.; Guo, S.; and Xu, W. 2022. Layer-wised model aggregation for personalized federated learning. In CVPR.
  29. 29.McMahan, B.; Moore, E.; Ramage, D.; Hampson, S.; and y Arcas, B. A. 2017. Communication-Efficient Learning of Deep Networks from Decentralized Data. In AISTATS.
  30. 30.Nilsback, M.-E.; and Zisserman, A. 2008. Automated flower classification over a large number of classes. In 2008 Sixth Indian conference on computer vision, graphics & image processing, 722–729. IEEE.
  31. 31.Pinheiro, P. O. 2018. Unsupervised domain adaptation with similarity learning. In CVPR.
  32. 32.Sandler, M.; Howard, A.; Zhu, M.; Zhmoginov, A.; and Chen, L.-C. 2018. Mobilenetv2: Inverted residuals and linear bottlenecks. In CVPR.
  33. 33.Schroff, F.; Kalenichenko, D.; and Philbin, J. 2015. Facenet: A unified embedding for face recognition and clustering. In CVPR.
  34. 34.Shamsian, A.; Navon, A.; Fetaya, E.; and Chechik, G. 2021. Personalized federated learning using hypernetworks. In ICML.
  35. 35.Shen, T.; Zhang, J.; Jia, X.; Zhang, F.; Huang, G.; Zhou, P.; Kuang, K.; Wu, F.; and Wu, C. 2020. Federated mutual learning. arXiv preprint arXiv:2006.16765.
  36. 36.Shin, K.; Kwak, H.; Kim, S. Y.; Ramstrom, M. N.; Jeong, J.; Ha, J.-W.; and Kim, K.-M. 2023. Scaling law for recommendation models: Towards general-purpose user representations. In AAAI.
  37. 37.Szegedy, C.; Liu, W.; Jia, Y.; Sermanet, P.; Reed, S.; Anguelov, D.; Erhan, D.; Vanhoucke, V.; and Rabinovich, A. 2015. Going deeper with convolutions. In CVPR.
  38. 38.T Dinh, C.; Tran, N.; and Nguyen, T. D. 2020. Personalized Federated Learning with Moreau Envelopes. In NeurIPS.
  39. 39.Tan, A. Z.; Yu, H.; Cui, L.; and Yang, Q. 2022a. Towards Personalized Federated Learning. IEEE Transactions on Neural Networks and Learning Systems. Early Access.
  40. 40.Tan, Y.; Long, G.; Liu, L.; Zhou, T.; Lu, Q.; Jiang, J.; and Zhang, C. 2022b. Fedproto: Federated Prototype Learning across Heterogeneous Clients. In AAAI.
  41. 41.Tan, Y.; Long, G.; Ma, J.; Liu, L.; Zhou, T.; and Jiang, J. 2022c. Federated Learning from Pre-Trained Models: A Contrastive Learning Approach. arXiv preprint arXiv:2209.10083.
  42. 42.Tanwisuth, K.; Fan, X.; Zheng, H.; Zhang, S.; Zhang, H.; Chen, B.; and Zhou, M. 2021. A prototype-oriented framework for unsupervised domain adaptation. NeurIPS.
  43. 43.Wang, H.; Yurochkin, M.; Sun, Y.; Papailiopoulos, D.; and Khazaeni, Y. 2020. Federated learning with matched averaging. arXiv preprint arXiv:2002.06440.
  44. 44.Wang, L.; Wang, M.; Zhang, D.; and Fu, H. 2023. Model Barrier: A Compact Un-Transferable Isolation Domain for Model Intellectual Property Protection. In CVPR.
  45. 45.Wen, D.; Jeon, K.-J.; and Huang, K. 2022. Federated dropout—A simple approach for enabling federated learning on resource constrained devices. IEEE wireless communications letters, 11(5): 923–927.
  46. 46.Wu, C.; Wu, F.; Lyu, L.; Huang, Y.; and Xie, X. 2022. Communication-efficient federated learning via knowledge distillation. Nature communications, 13(1): 2032.
  47. 47.Xu, W.; Xian, Y.; Wang, J.; Schiele, B.; and Akata, Z. 2020. Attribute prototype network for zero-shot learning. NeurIPS.
  48. 48.Yang, H.-M.; Zhang, X.-Y.; Yin, F.; and Liu, C.-L. 2018. Robust classification with convolutional prototype learning. In CVPR.
  49. 49.Yang, X.; Huang, W.; and Ye, M. 2023. Dynamic Personalized Federated Learning with Adaptive Differential Privacy. In NeurIPS.
  50. 50.Yi, L.; Wang, G.; Liu, X.; Shi, Z.; and Yu, H. 2023. FedGH: Heterogeneous Federated Learning with Generalized Global Header. arXiv preprint arXiv:2303.13137.
  51. 51.Yu, Q.; Liu, Y.; Wang, Y.; Xu, K.; and Liu, J. 2022. Multimodal Federated Learning via Contrastive Representation Ensemble. In ICLR.
  52. 52.Zhang, J.; Gu, Z.; Jang, J.; Wu, H.; Stoecklin, M. P.; Huang, H.; and Molloy, I. 2018a. Protecting intellectual property of deep neural networks with watermarking. In ASIA-CCS.
  53. 53.Zhang, J.; Guo, S.; Guo, J.; Zeng, D.; Zhou, J.; and Zomaya, A. 2023a. Towards Data-Independent Knowledge Transfer in Model-Heterogeneous Federated Learning. IEEE Transactions on Computers.
  54. 54.Zhang, J.; Guo, S.; Ma, X.; Wang, H.; Xu, W.; and Wu, F. 2021. Parameterized Knowledge Transfer for Personalized Federated Learning. In NeurIPS.
  55. 55.Zhang, J.; Hua, Y.; Cao, J.; Wang, H.; Song, T.; XUE, Z.; Ma, R.; and Guan, H. 2023b. Eliminating Domain Bias for Federated Learning in Representation Space. In NeurIPS.
  56. 56.Zhang, J.; Hua, Y.; Wang, H.; Song, T.; Xue, Z.; Ma, R.; Cao, J.; and Guan, H. 2023c. GPFL: Simultaneously Learning Global and Personalized Feature Information for Personalized Federated Learning. In ICCV.
  57. 57.Zhang, J.; Hua, Y.; Wang, H.; Song, T.; Xue, Z.; Ma, R.; and Guan, H. 2023d. FedALA: Adaptive Local Aggregation for Personalized Federated Learning. In AAAI.
  58. 58.Zhang, J.; Hua, Y.; Wang, H.; Song, T.; Xue, Z.; Ma, R.; and Guan, H. 2023e. FedCP: Separating Feature Information for Personalized Federated Learning via Conditional Policy. In KDD.
  59. 59.Zhang, K.; and Sato, Y. 2023. Semantic Image Segmentation by Dynamic Discriminative Prototypes. IEEE Transactions on Multimedia.
  60. 60.Zhang, L.; Shen, L.; Ding, L.; Tao, D.; and Duan, L.-Y. 2022. Fine-Tuning Global Model Via Data-Free Knowledge Distillation for Non-IID Federated Learning. In CVPR.
  61. 61.Zhang, Y.; Xiang, T.; Hospedales, T. M.; and Lu, H. 2018b. Deep mutual learning. In CVPR.
  62. 62.Zhong, Z.; Li, J.; Ma, L.; Jiang, H.; and Zhao, H. 2017. Deep residual networks for hyperspectral image classification. In IEEE international geoscience and remote sensing symposium (IGARSS).
  63. 63.Zhu, Z.; Hong, J.; and Zhou, J. 2021. Data-Free Knowledge Distillation for Heterogeneous Federated Learning. In ICML.
  64. 64.Zhuang, W.; Chen, C.; and Lyu, L. 2023. When Foundation Model Meets Federated Learning: Motivations, Challenges, and Future Directions. arXiv preprint arXiv:2306.15546.

Citation

MLA
Zhang, J., et al. “FedTGP: Trainable Global Prototypes with Adaptive-Margin-Enhanced Contrastive Learning for Data and Model Heterogeneity in Federated Learning”. arXiv, 2024, http://arxiv.org/abs/2401.03230v1.
APA
Zhang, J., Liu, Y., Hua, Y., & Cao, J. (2024). FedTGP: Trainable Global Prototypes with Adaptive-Margin-Enhanced Contrastive Learning for Data and Model Heterogeneity in Federated Learning. arXiv. http://arxiv.org/abs/2401.03230v1
Chicago
Zhang, J., Y. Liu, Y. Hua, and J. Cao. 2024. “FedTGP: Trainable Global Prototypes with Adaptive-Margin-Enhanced Contrastive Learning for Data and Model Heterogeneity in Federated Learning”. arXiv. http://arxiv.org/abs/2401.03230v1.
Harvard
Zhang, J. et al. (2024) “FedTGP: Trainable Global Prototypes with Adaptive-Margin-Enhanced Contrastive Learning for Data and Model Heterogeneity in Federated Learning”, arXiv [Preprint]. Available at: http://arxiv.org/abs/2401.03230v1.
Vancouver
1. Zhang J, Liu Y, Hua Y, Cao J (2024) FedTGP: Trainable Global Prototypes with Adaptive-Margin-Enhanced Contrastive Learning for Data and Model Heterogeneity in Federated Learning. arXiv

BibTeX

@article{zhang2024fedtgp,
  title = {FedTGP: Trainable Global Prototypes with Adaptive-Margin-Enhanced Contrastive Learning for Data and Model Heterogeneity in Federated Learning},
  author = {Zhang, Jianqing and Liu, Yang and Hua, Yang and Cao, Jian},
  year = {2024},
  journal = {arXiv},
  url = {http://arxiv.org/abs/2401.03230v1},
  eprint = {2401.03230}
}
Metadata:arXiv

Source Code

This paper has an official code repository available. Click below to access the source code.

View Repository

Access the Paper

This paper is available from its original source. Click below to access the PDF.

Open PDF