Federal learning method and device based on global personalized aggregation
By decomposing the model into feature extractors and classifiers, and performing global personalized aggregation in federated learning, and dynamically adjusting the weights to minimize variance and deviations, the model performance and convergence problems in data heterogeneous scenarios are solved, and the optimization of the personalized model and overall performance improvement are achieved.
Patent Information
- Application Number
- CN202510323634.9
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-03-19
- Publication Date
- 2025-07-04
AI Technical Summary
Traditional federated learning methods have poor model performance and convergence when facing data heterogeneous scenarios, and lack personalized solutions, resulting in poor performance of client models and slow convergence speed. The existing personalized federated learning methods have failed to effectively combine client and server-side optimization.
The model is decomposed into feature extractors and classifiers, and by performing global personalized aggregation in each round of communication, the client transmits feature extractors and classifier parameters, and dynamically adjusts feature extractor weights by calculating prototype similarity, and optimizes classifier weights to solve the quadratic planning problem to minimize variance and deviation.
In data heterogeneous scenarios, the model performance and applicability are improved, ensuring that each client obtains a personalized model, and improving the generalization ability and adaptability of the overall model.
Smart Images

Figure CN120258172A_ABST
Abstract
Description
Technical Field
[0001] This application relates to the technical fields of artificial intelligence, computer vision, and federated learning, and particularly relates to a federated learning method and device based on global personalized aggregation. Background Art
[0002] Contemporary learning tasks mainly rely on deep neural networks (DNNs), which require a large amount of training data to achieve satisfactory model performance. However, in practical applications, data is usually scattered among different participating parties. With the increasing concern about privacy issues and the implementation of relevant regulations, the challenges and costs of directly obtaining and centralizing data for learning and training purposes are also increasing. To address the above challenges, Federated Learning (FL), as a promising machine learning method, allows clients to aggregate by sharing local model parameters without accessing each other's data, thereby collaboratively training a global model. The Federated Averaging (FedAvg) algorithm is a classic federated learning method. In each round of communication, FedAvg selects a portion of clients for local training, and then aggregates their parameters on the server according to the number of samples owned by the client as the weight. Federated learning algorithms have been widely applied in various practical scenarios, including medical imaging, object detection, and predicting the next word of a user.
[0003] However, traditional federated learning methods face several key challenges: (1) poor model performance and convergence in the face of highly data heterogeneous scenarios; (2) lack of personalized solutions. For example, the next word prediction model trained for users using FedAvg may not always be effective because users have different habits, and this single global model may deviate significantly from the individual's optimal model. In addition, studies have shown that data heterogeneity can cause a "drift" phenomenon in the parameter updates of clients, resulting in poor performance and slow convergence speed. The existence of these problems will reduce the performance of the global model on the client side and may even make the affected clients reluctant to participate in the federated learning process. To address these problems, scholars have proposed Personalized Federated Learning (PFL) solutions.
[0004] Currently, there are mainly two methods: (1) adjusting the local model on the client side; (2) optimizing the global model on the server side. In the first method, SCAFFOLD uses control variables to solve the client "drift" problem; MOON uses contrastive learning to constrain the training process of the client local model. Specifically, it guides the update direction of the local model by comparing the representations between the local model, the global model, and the historical model, making it closer to the global model and farther from the historical model. In contrast, in the second method, FedDF and FedFTG further optimize the global model by introducing additional fine-tuning steps to generate personalized models; while FedPAC divides the clients into feature extractors and classifiers, and designs personalized classifier aggregation weights for each client to better adapt to the unique data distribution of each client. Although there are numerous methods for personalized federated learning currently, few works consider adjusting the local model on the client side and optimizing the aggregation weights on the server side simultaneously. Summary of the Invention
[0005] This application aims to solve at least one of the technical problems in the related art to some extent.
[0006] To this end, the first object of this application is to propose a federated learning method based on global personalized aggregation, which solves the data heterogeneity problem in the existing methods, enables clients to perform collaborative training in a personalized manner, ensures that each participant can benefit from it, improves the performance and applicability of the overall model, and can better adapt to the requirements of different application scenarios, thereby improving the overall efficiency of the federated learning system.
[0007] The second object of this application is to propose a federated learning device based on global personalized aggregation.
[0008] To achieve the above object, the first aspect embodiment of this application proposes a federated learning method based on global personalized aggregation, which includes: decomposing the model into a feature extractor and a classifier, and in each round of communication, simultaneously performing global personalized aggregation of the feature extractor and the classifier. Among them, in the global personalized aggregation of the feature extractor, each client transmits its own feature extractor parameters and local prototypes in each round of communication, and dynamically adjusts the feature extractor weights of each client by calculating the prototype similarity between clients; in the global personalized aggregation of the classifier, each client transmits its own classifier parameters in each round of communication, and optimizes the classifier weights of each client by solving a quadratic programming problem to minimize variance and bias.
[0009] To achieve the above object, the second aspect embodiment of the present invention proposes a federated learning device based on global personalized aggregation, including clients and a server, and this device implements the above-mentioned federated learning method based on global personalized aggregation.
[0010] The federated learning method and device based on global personalized aggregation in the embodiments of the present application propose a global personalized aggregation strategy, which decomposes the model into a feature extractor and a classifier, and performs personalized aggregation separately. In the global personalized aggregation process of the feature extractor, client i and client j transmit their respective feature extractor parameters f i and f j , as well as local prototypes C i and C j in each round of communication. By calculating the prototype similarity between clients, the aggregation weights are dynamically adjusted to ensure that each client can obtain a personalized feature extractor suitable for its data distribution; in the global personalized aggregation process of the classifier, client i and client j transmit their respective classifier parameters c i and c j . By solving a quadratic programming problem, the aggregation weights of the classifier are optimized to minimize variance and bias, ensuring that the classifier can better adapt to the local data distribution during the personalized aggregation process. Through the proposed global personalized aggregation strategy for federated learning, this embodiment can effectively improve the model performance in a data heterogeneous scenario, ensure that each client can obtain a personalized model from federated learning, and thus improve the applicability and generalization ability of the overall model.
[0011] Some of the additional aspects and advantages of the present application will be given in the following description, some will become obvious from the following description, or will be understood through the practice of the present application. Description of the Drawings
[0012] The above and / or additional aspects and advantages of the present application will become obvious and easy to understand from the following description of the embodiments in conjunction with the drawings, where:
[0013] Figure 1 is a schematic flowchart of the image classification method based on global personalized aggregation in the federated learning scenario of the embodiments of the present application. Detailed Embodiments
[0014] The embodiments of the present application will be described in detail below. The examples of the embodiments are shown in the drawings, where the same or similar reference numerals denote the same or similar elements or elements having the same or similar functions throughout. The embodiments described below with reference to the drawings are exemplary and are intended to explain the present application and should not be construed as limiting the present application.
[0015] The federated learning method and device based on global personalized aggregation in the embodiments of the present application will be described below with reference to the drawings.
[0016] The federated learning method based on global personalized aggregation includes the following steps:
[0017] Step 101: Decompose the model into a feature extractor and a classifier. In each round of communication, global personalized aggregation of both the feature extractor and the classifier is performed. Among them,
[0018] In the global personalized aggregation of the feature extractor, each client transmits its own feature extractor parameters and local prototypes in each round of communication. By calculating the prototype similarity between clients, the feature extractor weights of each client are dynamically adjusted;
[0019] In the global personalized aggregation of the classifier, each client transmits its own classifier parameters in each round of communication. By solving a quadratic programming problem to minimize variance and bias, the classifier weights of each client are optimized.
[0020] Furthermore, in the embodiments of this application, the federated learning task is defined as:
[0021] Suppose there are m clients and a central server. All clients communicate with the server under data protection requirements to collaboratively train personalized models. Each client i has a private dataset with a data distribution of P i (x, y), where x is the input feature and y is the corresponding label, and the labels are divided into K categories. There are significant differences in both the label distribution and the number of samples between clients i and j. In this embodiment, g is defined as the combination of the feature extractor f and the classifier c, where f is parameterized by the parameter space to map the input feature x to a latent vector; c is parameterized by the parameter space to map the latent vector to the output result. Formally, this can be expressed as and Let be the loss function of client i, which is evaluated on the data samples (x, y) drawn from its distribution P i (x, y). The global optimization objective of personalized federated learning can be expressed as:
[0022]
[0023] where W = {w1, w2,..., w m} represents the set of model parameters of all clients. In actual situations, the true data distribution P i (x, y) is unknown, so the empirical risk minimization method is adopted. Suppose client i samples n i data from its distribution, denoted as representing the empirical distribution of P i . The empirical training objective of the personalized model can be re-expressed as:
[0024]
[0025] Sub-item represents the empirical loss calculated on the local dataset D i . The sub-item represents the regularization term used to avoid overfitting, where Ω represents some global or local constraint, such as the L2 regularization that penalizes large weights.
[0026] Traditional federated learning (FL) aims to find the shared optimal global model w by solving the empirical training objective formula of the personalized model * . However, in a heterogeneous setting, this embodiment focuses more on optimizing the personalized model for each client. This embodiment hopes to find the optimal weights to optimize the objective of the client and minimize F(W). Here represents the optimal model of the i-th client. In an independent and identically distributed (IID) data setting, the optimal models of any clients are very close, that is However, in a non-independent and identically distributed (non-IID) data setting, a general global model cannot achieve the best performance on all clients because the optimal solutions of different clients are usually inconsistent.
[0027] Specifically, in the embodiments of this application, a regularization term is introduced in federated learning to achieve "global-local" alignment based on prototypes, including:
[0028] To enable the local model to more effectively utilize local and global information, this embodiment decomposes the model into two independent parts: a feature extractor and a classifier. The feature extractor is implemented through convolutional layers and is used to learn the representations of samples; while the classifier is implemented through fully connected layers to generate the final classification vector. The feature embedding function is parameterized by the parameter w f where d is the dimension of the feature embedding. This embodiment represents z = f(w f ; x) as the embedding of x. The classification function c: Z → R d is parameterized by the parameter w c where d is the dimension of z. This embodiment represents as the final classification vector. In the case of data heterogeneity, the local label distribution of each client is skewed and the data volume is insufficient, which causes the local training process to gradually deviate from the global optimal solution. Therefore, this embodiment introduces a regularization term to constrain the distance between the local objective and the global objective. The present invention uses the feature embedding vector output by the feature extractor as the basis for calculating the regularization term.
[0029] For the i-th client, the local prototype represents the mean of the embedding vectors of the k-th class:
[0030]
[0031] where D i,k is a subset containing the k-th class of training data, which is a part of the local dataset D i of represents the embedding vector of x.
[0032] For the k-th class, the global prototype represents the mean of all local prototypes belonging to the k-th class:
[0033]
[0034] where is the embedding vector of the local prototype matrix of the k-th class of the i-th client. |D i,k | represents the number of samples of the k-th class held by the i-th client.
[0035] This embodiment uses the local prototype matrix and the global prototype matrix to construct a regularization term, which can introduce global information in the local training of the model. The local regularization term is defined as:
[0036]
[0037] where, is the embedding vector of the local prototype of the k-th class of the i-th client, is the corresponding global prototype of the k-th class. λ is an adjustable hyperparameter used to balance the supervised loss and the regularization loss.
[0038] Furthermore, in the embodiment of the present application, during federated learning, global personalized aggregation is performed, including:
[0039] In the context of data heterogeneity, each client will encounter challenges of insufficient data and unbalanced data distribution. To continuously improve the model performance, this embodiment adopts personalized aggregation weights for each client in each communication round to obtain a personalized aggregation model. This is different from the practice of using a unified global aggregation model in traditional federated learning methods. In addition, different from the previous federated learning process, this embodiment simultaneously performs personalized aggregation of the classifier and the feature extractor. This embodiment assigns personalized aggregation weights to each client through a personalized aggregation strategy, optimizing the aggregation process of the feature extractor and the classifier. This strategy dynamically adjusts the aggregation weights by calculating the prototype similarity and the number of samples between clients, ensuring that each client can obtain a personalized model suitable for its data distribution.
[0040] Specifically, in the (t + 1)-th global communication round, the aggregation functions of the feature extractor and the classifier can be defined as:
[0041]
[0042] Among them, m is the number of clients, and the subscript i represents the i-th client. is the parameter space of the feature extractor f. is the parameter space of the classifier c, and α ij represents the weighting factor of the j-th client when the i-th client aggregates the feature extractor, and β ij represents the weighting factor of the j-th client when the i-th client aggregates the classifier.
[0043] Specifically, in the embodiments of the present application, the calculation process of the aggregation weight of the feature extractor includes:
[0044] When performing federated learning, during global aggregation, the feature extractor should preferentially capture global information to improve feature representation and thus enhance the generalization ability of the model. Therefore, a more uniform weight distribution may be beneficial. If each client can selectively focus on clients with similar representations during the global aggregation process, it will continuously contribute to enhancing the performance of the model. Then the key issue becomes: effectively evaluating the similarity between clients. For this purpose, in this embodiment, the similarity between two clients is measured by calculating the prototype matrix between the two clients. Let and respectively represent the local prototypes of the i-th client and the j-th client for class k. Then, the distance between the i-th client and the j-th client can be expressed as:
[0045]
[0046] Here, |D i | represents the total number of samples owned by the i-th client, while |D i,k | represents the total number of samples of the i-th client under class k. The variable K represents the total number of classes. is the reciprocal of the distance, which converts the distance P i,j into a weight. In addition, the number of samples owned by each client is also a key consideration. Therefore, this embodiment will consider both the prototype matrix similarity and the number of samples owned by the client. The formula for calculating the weight factor from client i to client j is as follows:
[0047]
[0048] where μ is a hyperparameter used to balance the influence of prototype similarity and the number of samples on α i,j , and α i,j is obtained by normalizing .
[0049] Furthermore, in the embodiments of the present application, the process of calculating the classifier aggregation weights includes:
[0050] Compared with the feature extractor, since the classifier directly affects the final output, this embodiment believes that it should pay more attention to learning from local samples, so that the personalized weights can play a more important role. This embodiment uses the classifiers of other clients to reduce the variance, but this will increase the deviation between the result and the true distribution P i (x, y). Therefore, this embodiment solves a quadratic programming problem to find an optimal classifier weight distribution that can minimize the variance and deviation. By solving the quadratic programming problem, this embodiment optimizes the aggregation weights of the classifier to minimize the variance and deviation, ensuring that the classifier can better adapt to the local data distribution during the personalized aggregation process.
[0051] Specifically, this embodiment is based on the distance matrix between clients (the distance matrix P obtained by the above solution) and the within-class variance to meet this requirement. These formulas are shown as follows:
[0052]
[0053] where p k represents the probability of class k, and F k is the prototype matrix composed of all the hidden layer features of class k. Therefore can also be expressed as:
[0054]
[0055] Then, the optimization problem can be formulated as:
[0056]
[0057]
[0058] where β is the optimal classifier weight to be solved.
[0059] In the federated learning method based on global personalized aggregation of the embodiments of the present application, client i and client j respectively use local data x i and x j for training. The feature extractors f i and f j map the input features to the latent vectors Z i and Z j , and the classifiers c i and c j then map the latent vectors to the output results y i and y j。This embodiment proposes a global personalized aggregation strategy, which decomposes the model into a feature extractor and a classifier, and performs personalized aggregation separately. In the process of global personalized aggregation of the feature extractor, client i and client j transmit their respective feature extractor parameters f i and f j , as well as local prototypes C i and C j in each round of communication. By calculating the prototype similarity between clients, the aggregation weights are dynamically adjusted to ensure that each client can obtain a personalized feature extractor suitable for its data distribution; in the process of global personalized aggregation of the classifier, client i and client j transmit their respective classifier parameters c i and c j . By solving a quadratic programming problem, the aggregation weights of the classifier are optimized to minimize variance and bias, ensuring that the classifier can better adapt to the local data distribution during the personalized aggregation process. This embodiment performs federated learning through the proposed global personalized aggregation strategy, which can effectively improve the model performance in the data heterogeneous scenario, ensure that each client can obtain a personalized model from federated learning, and thus improve the applicability and generalization ability of the overall model.
[0060] Figure 1 FIG. is a schematic flow chart of an image classification method based on global personalized aggregation in a federated learning scenario. In this method, each client has its own image data, extracts image features through a feature extractor, and realizes image classification through a classifier.
[0061] To implement the above embodiment, the present application also proposes a federated learning device based on global personalized aggregation, which implements the above federated learning method based on global personalized aggregation.
[0062] In the description of this specification, the description with reference to the terms "one embodiment", "some embodiments", "example", "specific example" or "some examples", etc. means that the specific features, structures, materials or characteristics described in connection with the embodiment or example are included in at least one embodiment or example of the present application. In this specification, the schematic representations of the above terms do not necessarily refer to the same embodiment or example. Moreover, the specific features, structures, materials or characteristics described can be combined in any one or more embodiments or examples in a suitable manner. In addition, without contradiction, those skilled in the art can combine and combine the different embodiments or examples described in this specification and the features of different embodiments or examples.
[0063] In addition, the terms "first" and "second" are used for descriptive purposes only and should not be construed as indicating or implying relative importance or implicitly specifying the quantity of the technical features indicated. Thus, features defined with "first" and "second" may explicitly or implicitly include at least one such feature. In the description of the present application, "a plurality of" means at least two, such as two, three, etc., unless otherwise specifically defined.
[0064] Any process or method description represented in a flowchart or otherwise described herein may be understood to represent a module, segment, or portion of code including one or more executable instructions for implementing a customized logical function or process. The scope of the preferred embodiments of the present application includes additional implementations, where functions may be executed in a substantially simultaneous manner or in a reverse order according to the functions involved, rather than in the order shown or discussed, which should be understood by those skilled in the art to which the embodiments of the present application pertain.
[0065] The logic and / or steps represented in a flowchart or otherwise described herein, for example, may be considered a sequenced list of executable instructions for implementing a logical function, and may be embodied specifically in any computer-readable medium for use by or in connection with an instruction execution system, apparatus, or device, such as a computer-based system, a system including a processor, or other systems that can fetch and execute instructions from the instruction execution system, apparatus, or device. For the purposes of this specification, a "computer-readable medium" can be any device that can contain, store, communicate, propagate, or transport a program for use by or in connection with an instruction execution system, apparatus, or device. More specific examples (a non-exhaustive list) of the computer-readable medium include the following: an electrical connection portion having one or more wirings (electronic device), a portable computer diskette (magnetic device), a random access memory (RAM), a read-only memory (ROM), an erasable programmable read-only memory (EPROM or flash memory), an optical fiber device, and a portable compact disc read-only memory (CDROM). Additionally, the computer-readable medium can even be paper or other suitable media on which the program can be printed, as the program can be obtained electronically, for example, by optically scanning the paper or other media, followed by editing, interpretation, or otherwise appropriate processing if necessary, and then stored in a computer memory.
[0066] It should be understood that each part of the present application can be implemented by hardware, software, firmware, or a combination thereof. In the above embodiments, multiple steps or methods can be implemented by software or firmware stored in a memory and executed by a suitable instruction execution system. For example, if implemented by hardware, as in another embodiment, any one of the following techniques known in the art or a combination thereof can be used: discrete logic circuits with logic gate circuits for implementing logic functions on data signals, application specific integrated circuits with appropriate combinational logic gate circuits, programmable gate arrays (PGAs), field programmable gate arrays (FPGAs), etc.
[0067] Those of ordinary skill in the art can understand that all or part of the steps carried by the method of implementing the above embodiments can be completed by instructing relevant hardware through a program. The said program can be stored in a computer-readable storage medium. When the program is executed, it includes one or a combination of the steps of the method embodiments.
[0068] In addition, in each embodiment of the present application, each functional unit can be integrated into a processing module, or each unit can exist physically alone, or two or more units can be integrated into one module. The above integrated module can be implemented in the form of hardware or in the form of a software functional module. When the above integrated module is implemented in the form of a software functional module and sold or used as an independent product, it can also be stored in a computer-readable storage medium.
[0069] The above-mentioned storage medium can be a read-only memory, a magnetic disk, an optical disk, etc. Although the embodiments of the present application have been shown and described above, it can be understood that the above embodiments are exemplary and should not be construed as limiting the present application. Those of ordinary skill in the art can make changes, modifications, substitutions, and variations to the above embodiments within the scope of the present application.
Claims
1. A federated learning method based on global personalized aggregation, characterized in that, Including: Decompose the model into a feature extractor and a classifier. In each round of communication, simultaneously perform global personalized aggregation of the feature extractor and the classifier. Wherein, In the global personalized aggregation of the feature extractor, each client transmits its own feature extractor parameters and local prototypes in each round of communication. By calculating the prototype similarity between clients, dynamically adjust the weights of the feature extractor for each client. In the global personalized aggregation of the classifier, each client transmits its own classifier parameters in each round of communication. By solving a quadratic programming problem, minimize the variance and bias to optimize the classifier weights for each client.
2. The method according to claim 1, characterized in that The method further includes: Set the global optimization objective function of personalized federated learning as: where m is the number of clients, and the subscript i represents the i-th client. denotes the empirical loss calculated on the local dataset D i where D i is sampled from the data distribution P i (x, y) of the private dataset of client i to obtain n i data points, where x is the input feature and y is the class label, is the parameter space of the feature extractor f, is the parameter space of the classifier c, is the regularization term used to avoid overfitting, and Ω represents the global or local constraint. The goal of the personalized federated learning is: by solving the global optimization objective function, find the optimal weights which is the optimal model for the m-th client to achieve the optimization goal of the client.
3. The method according to claim 1, characterized in that The method further includes: During model training, introduce a regularization term to constrain the distance between the local objective and the global objective. Wherein, the regularization term is constructed using the local prototype matrix and the global prototype matrix. The regularization term introduces global information in local model training. The regularization term is expressed as: where the subscript i represents the i-th client, and the local dataset D i is sampled from the data distribution P i of the private dataset of client i with n i data points obtained from where x is the input feature and y is the class label, is the parameter space of the feature extractor f, is the global prototype, and λ is an adjustable hyperparameter, is the embedding vector of the local prototype matrix of the k-th class of the i-th client, denoted as: Among them, D i,k is a subset of D i that contains the k-th type of training data, and f(w i f ; x) is the embedding vector of x are all the global prototypes belonging to the k-th class, expressed as:
4. The method according to claim 1, wherein In each round of communication, simultaneously perform global personalized aggregation of the feature extractor and the classifier, including: In the (t + 1)-th global communication round, define the aggregation functions of the feature extractor and the classifier respectively as: where m is the number of clients, and the subscript i represents the i-th client. is the parameter space of the feature extractor f. is the parameter space of the classifier c, and α ij represents the weighting factor of the j-th client when the i-th client aggregates the feature extractor, and β ij represents the weighting factor of the j-th client when the i-th client aggregates the classifier.
5. The method according to claim 4, wherein α ij The calculation formula is as follows: Among them, α ij is obtained from through normalization, represents the weight from client i to client j, expressed as: μ is a hyperparameter, and the local dataset D i is sampled from the data distribution P i of the private dataset of client i to obtain n i data points from (x, y), where x is the input feature and y is the class label is the reciprocal of the distance, indicating the conversion of the distance into a weight, which is expressed as: P ij is the distance between the i-th client and the j-th client, expressed as: Among them, D i,k is a subset of D i that contains the subset of the k-th type of training data, where K represents K categories, is the embedded vector of the local prototype matrix of the k-th type of the i-th client.
6. The method according to claim 4, wherein Said β ij The calculation process includes: By solving a quadratic programming problem, find the optimal classifier weights that minimize the variance and bias. Wherein, The quadratic programming problem is: Among them, β is the optimal classifier weight to be solved, and P is the distance matrix between clients, is the within-class variance, Local dataset D i is sampled from the data distribution P of the private dataset of client i i n data from (x, y) i to obtain,[[]] where x is the input feature, y is the class label, K represents K classes, and p k represents the probability of class k, and D i,k is D i is the subset of the training data of the k-th class included in D k is the prototype matrix composed of all hidden layer features of class k, which is expressed as:
7. A federated learning device based on global personalized aggregation, characterized in that, Including clients and a server. The device implements the federated learning method based on global personalized aggregation as described in claims 1-6.