Personalized federal learning method based on trainable prototype
By introducing the FedETP algorithm of contrast learning and adversarial training in federated learning, the prototype offset problem in the environment of model heterogeneity and data heterogeneity is solved, the global prototype quality and classification accuracy are improved, communication overhead is reduced, and privacy and intellectual property rights are protected.
Patent Information
- Application Number
- CN202510285811.9
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-03-11
- Publication Date
- 2025-06-27
AI Technical Summary
The existing federated learning methods have prototype offset problems in the environment of model heterogeneity and data heterogeneity, resulting in unclear boundaries of global prototype separation, reduced classification accuracy, large communication overhead, and difficulty in protecting privacy and intellectual property rights.
A personalized federated learning algorithm (FedETP) based on trainable prototypes is proposed. Through comparative learning and adversarial training, the training method of global prototype generator is improved, the representativeness of local prototypes is enhanced, and a collaborative working mechanism is designed between the server and the client to reduce communication overhead.
It effectively improves the global prototype quality in the environment of model heterogeneity and data heterogeneity, enhances inter-class separation, improves classification accuracy, reduces communication overhead, and protects client data privacy and model intellectual property rights.
Smart Images

Figure CN120218187A_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the technical field of federated learning, and particularly to a method and system for realizing personalized federated learning based on trainable class prototypes in a model heterogeneous environment. Technical Background
[0002] In the field of federated learning, the performance of traditional methods drops significantly when facing statistical heterogeneity. Although personalized federated learning algorithms aim to solve this problem by learning personalized model parameters, most existing algorithms still assume that all client model architectures are the same, and there are many problems in the training process. For example, transmitting client model parameters brings non-negligible communication overhead, and there are huge challenges in privacy and intellectual property protection.
[0003] To solve the above problems, heterogeneous federated learning has emerged, which allows clients to have different model structures and Non-IID data. However, existing heterogeneous federated learning methods also have defects. Some knowledge distillation-based methods highly rely on the quality of the public dataset, and data-free knowledge distillation methods need to share auxiliary models, increasing the communication overhead. Although the prototype-based heterogeneous federated learning algorithm uses prototypes composed of lightweight class representations to reduce the communication overhead, when aggregating prototypes, due to model heterogeneity and data heterogeneity, the prototypes generated by clients are inconsistent in scale and separation boundaries, resulting in unclear separation boundaries of the globally generated prototype by weighted average, reducing the representativeness of the global prototype for various types of data. Existing methods also do not fully consider the negative impact of prototypes generated by poorly performing client models on the training of the global prototype, as well as problems such as insufficient utilization of the global prototype by local clients and communication redundancy. Summary of the Invention
[0004] The present invention proposes a personalized federated learning algorithm (FedETP) based on trainable class prototypes, aiming to improve the quality of global prototypes in an environment of model heterogeneity and data heterogeneity, enhance the separability between classes; optimize the client local model through contrastive learning and adversarial training to improve the classification accuracy; reduce the communication overhead, and protect the client data privacy and model intellectual property rights.
[0005] In terms of algorithm design, the present invention improves the way clients use the global prototype for training. It not only narrows the distance between the local class representation and the corresponding global class prototype, but also widens the distance from non-class global prototypes, enhancing the ability of the local classifier to distinguish different class samples. At the same time, the training method of the global prototype generator is improved, using contrastive learning to generate effective global prototypes, and using adversarial learning to enhance the ability to mine hard samples. In addition, the global prototype generator is sent to local clients for personalized training, so that the class prototypes generated by the local prototype generator can not only represent local clients, but also have the commonality of global prototypes.
[0006] In terms of system construction, a collaborative working mechanism between the server side and the client side is designed. The server side receives the local prototypes generated by the local clients, obtains the average global prototype through weighted averaging. The global prototype generator uses the client prototypes and the average global prototype for training, generates the global prototype and the global prototype generator, and distributes them to the clients. After receiving them, the clients use the global prototype and the locally generated prototype to perform enhanced adversarial training on the local feature extractor, and at the same time train the local prototype generator to generate local prototypes and upload them to the server.
[0007] Through the above design, the present invention can effectively handle the prototype deviation problem caused by model heterogeneity, and in multiple datasets and different model heterogeneity environments, the model performance is better than the mainstream federated learning algorithms in the heterogeneous environment based on prototypes. Brief Description of the Drawings
[0008] Figure 1 is the architecture diagram of the FedETP system, showing the interaction process between the client and the server;
[0009] Figure 2 is the schematic diagram of the client data distribution under the Non-IID data distribution;
[0010] Figure 3 is the client training process;
[0011] Figure 4 is the server training process;
[0012] Figure 5 is the accuracy change of FedETP and FedTGP on the Cifar10 dataset with the number of communication rounds;
[0013] Figure 6 is the algorithm experiment accuracy result on the Cifar10 and Cifar100 datasets under different model heterogeneities;
[0014] Figure 7 is the accuracy change of this algorithm compared with other algorithms under different numbers of clients;
[0015] Figure 8 is the ablation experiment result of FedETP. Detailed Implementation Manner
[0016] The following combines the attached drawings to illustrate the detailed implementation manner of the present invention. Figure 1 is the overall flow chart of the algorithm.
[0017] Step 1: System initialization
[0018] The server initializes the global prototype generator G. For each category c, according to the given input category vector, the generator outputs the initial global prototype
[0019] Each client configures a heterogeneous feature extractor \(w\) k and a classifier \(\theta\) k , for example, the feature extractor of client 1 is selected as ResNet18, the feature extractor of client 2 is selected as MobileNet_v2, and the classifier structures of client 1 and 2 are kept consistent. The local datasets of the clients are divided according to the Dirichlet distribution (\(\alpha = 0.1\)) to simulate the Non-IID data distribution. Each client only holds partial category data, and the number of samples in each category varies significantly. The distribution status of the divided datasets is as shown in Figure 2 shown.
[0020] Step 2: Local training on the client side
[0021] The client receives the global prototype set and the generator \(G\) sent by the server.
[0022] The client calculates the local prototype according to the categories held by the local dataset:
[0023]
[0024] where \(f\) k (x; \(w\) k ) is the feature vector output by the feature extractor, and \(D\) k,c is the sample subset of category \(c\) in client \(k\).
[0025] Train the local model: Use a dual-supervised loss function to optimize the feature extractor:
[0026]
[0027] where \(\varPhi\) is the Euclidean distance.
[0028] Update the model parameters by combining the classifier correction loss:
[0029]
[0030] The total loss function of the client is:
[0031]
[0032] where represents the training loss of the local client on the local data. \(t\) is the current communication round, \(T\) is the total number of rounds, and \(\lambda\) is a hyperparameter. Training the local model according to the total loss function of the client can improve the model performance of the local model.
[0033] After the local client model training is completed, the client uses the received global prototype generator \(G\) as the local prototype generator \(G\) k , and fine-tunes it according to the following loss function:
[0034]
[0035] Among them is the trained local prototype generated by the local prototype generator according to the label y. After being trained, the local prototype generator generates a personalized prototype and uploads it to the server. The training process at this stage is as Figure 3 shown
[0036] Step 3: Server global training
[0037] Aggregate the client prototypes. For each category c, collect the client prototypes held in the current communication round for this category Calculate the adaptive weight matrix A c Obtain the client adaptive weights, where each element in the matrix is calculated as follows
[0038]
[0039] Among them and i≠j. The diagonal elements of the matrix are all 0
[0040] For the matrix A c Sum by row to obtain the initial category weight vector constructed for each client which is expressed as follows
[0041]
[0042] Normalize the weight vector to get
[0043]
[0044] For each category c, calculate the weighted contrast loss
[0045]
[0046] Among them is the initial global prototype output by the global prototype generator. By setting the weight β i , the contribution of low-quality prototypes is reduced
[0047] Use the simple weighted average prototype mostly adopted in existing algorithms to enhance the global prototype through adversarial training. First, calculate the simple weighted average prototype
[0048]
[0049] N c is the total number of samples of all clients belonging to category c which represents the number of clients containing category c. Based on this, calculate the adversarial loss
[0050]
[0051] The global prototype forced to be output by the generator Keep a distance from the simple average prototype to avoid falling into local optima. By integrating these two loss function terms, the total training loss function of the global prototype generator can be obtained:
[0052]
[0053] After training, the generator outputs a high-quality global prototype and distributes it to the client. This part of the process is as Figure 4 shown
[0054] Step 4: Iterative optimization
[0055] Repeat the client local training and server global training, that is, Step 2 and Step 3, until the model converges. The model convergence judgment condition is set to terminate the training when the validation set accuracy has not improved for 100 consecutive rounds
[0056] Finally, a personalized model adapted to each client is obtained
[0057] Experimental results
[0058] Accuracy comparison:
[0059] In the experiment, in an environment with two different models, the accuracy changes of FedETP and FedTGP on the Cifar10 dataset with the number of communication rounds are as Figure 5 shown, indicating the advantage of FedETP in model accuracy
[0060] The comparison experiment results with other algorithms in different model heterogeneous environments on the simulated data heterogeneity dataset in Step 1 show that on the Cifar100 dataset, in an environment with 8 different model structures, the average accuracy of FedETP is 72.06%, which is 1.27% higher than that of FedTGP (70.79%), as Figure 6 shown
[0061] Robustness verification:
[0062] When the number of clients increases from 50 to 100, the accuracy of FedETP only drops by 2.31%, while that of FedTGP drops by 4.73%, as Figure 7 shown
[0063] Ablation experiment results:
[0064] As Figure 8As shown, the accuracy after removing each module has a significant decrease compared to the complete algorithm. The experimental results show that each module in the algorithm is indispensable, and integrating each module can establish an effective personalized federated learning algorithm for training client models in heterogeneous environments.
Claims
1. A personalized federated learning method based on trainable class prototypes, which enhances the quality of trainable prototypes, enhances the performance of each client model, and improves the accuracy of each model on the test set.
2. According to claim 1, a personalized federated learning method based on trainable class prototypes is characterized in that The following steps are involved: S1. System initialization: The server initializes the global prototype generator. For each category, the generator outputs the initial global prototype according to the given input category vector. Each client is configured with a heterogeneous feature extractor and classifier, and the client local data set is divided according to the Dirichlet distribution. S2. Client local training: The client uses the local data set to calculate the local prototype and uses the global prototype sent by the server to train the local classifier. After the local client model is trained, the client uses the received global prototype generator as the local prototype generator and fine-tunes it using the local client model. After the local prototype generator is trained, a personalized prototype is generated and uploaded to the server. S3, server global training: Aggregate client prototypes, for each category, collect the client prototypes of the category in the current communication round, and calculate the category weight vector. The server-side generator uses the category weight vector to calculate the weighted contrast loss for each category to train the global prototype. In addition, the simple weighted average prototype is used to enhance the global prototype through adversarial training. First, the simple weighted average prototype is calculated, and then the adversarial loss is calculated. After training, the generator outputs a high-quality global prototype and sends it to the client. S4, iterative optimization: Repeat the client local training and server global training, i.e. S2 and S3, until the model converges. Finally, a personalized model adapted to each client is obtained.
3. The personalized federated learning method based on trainable class prototypes according to claim 2, characterized in that: The client local training in step S2 includes feature extractor training and classifier correction module: calculating the Euclidean distance between the local prototype and the global prototype; The global prototype is used as input to the client classifier to calculate the cross entropy loss.
4. The personalized federated learning method based on trainable class prototypes according to claim 2, characterized in that: The specific implementation of the adversarial training in step S3 includes: the adversarial loss is the Euclidean distance between the global prototype generated by the calculation generator and the simple weighted average global prototype.
5. The personalized federated learning method based on trainable class prototypes according to claim 2, characterized in that: The step of constructing the category weight vector in step S3 includes: calculating the distance matrix between client prototypes, generating an initial weight vector by summing up the rows, and normalizing it to obtain the final weight.
6. The personalized federated learning method based on trainable class prototypes according to claim 2, characterized in that: The network structure of the global prototype generator in step S3 is two fully connected layers, the input dimension is consistent with the feature space dimension, and the output dimension is the number of categories.
7. The personalized federated learning method based on trainable class prototypes according to claim 2, characterized in that: The model convergence judgment condition in step S4 is set to determine that the model converges if the accuracy of the test set does not improve for 100 consecutive rounds, and the training is terminated.
8. A personalized federated learning method based on a trainable class prototype according to any one of claims 1 to 7, characterized in that: The system includes: Server module, used for training and issuing global prototype generators; The client module configures heterogeneous feature extractors and classifiers, performs local prototype calculations and model optimization; The communication module is responsible for the transmission between the prototype and the generator. The communication content only contains the prototype and generator parameters. The final result is an improvement in the accuracy of the client model.
Citation Information
Cited By
Personalized global prototype federal learning method and system based on adaptive feature alignment
CN120509507A