Personalized federal learning method based on multi-teacher attention integrated distillation
Through the personalized federated learning method of multi-teacher attention integrated distillation, knowledge distillation is performed using client model output with high similarity, which solves the problems of poor convergence and model isomorphic limitation of traditional federated learning on heterogeneous data, and achieves the improvement of personalized learning.
Patent Information
- Application Number
- CN202510483747.5
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-04-17
- Publication Date
- 2025-07-18
AI Technical Summary
Traditional federated learning has poor convergence on heterogeneous data, lacks personalized solutions, and model isomorphic limitations are difficult to meet the needs of diversified equipment.
A personalized federal learning method of multi-teacher attention integrated distillation is adopted. Through the server-side attention mechanism and comparison learning technology, knowledge distillation is used to use the client model output with high similarity to generate a personalized model.
It significantly improves the personalized learning performance in heterogeneous scenarios, breaks the limitations of model isomorphism, and adapts to the diversified needs of different clients.
Smart Images

Figure CN120338053A_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the technical field of federated learning, and in particular to a personalized federated learning method based on multi-teacher attention integrated distillation. Background Art
[0002] With the development of artificial intelligence technology and the advent of the big data era, the rapid growth of data and the popularity of distributed storage have made this vast amount of information usually distributed across various devices. Traditional centralized machine learning converges information to a central server, which then processes, analyzes, and trains the model with the data. However, the process of uploading massive data consumes a large amount of communication resources. Especially in scenarios such as edge computing and the Internet of Things, the communication cost between devices is more significant.
[0003] To address these issues, federated learning has emerged. Federated learning is a distributed machine learning model where the central server distributes the initial model to all edge clients. Each client trains the model on its own data on the local device and then aggregates the model update information to the central server. However, there are several basic challenges in the general federated learning method: 1. Poor convergence on highly heterogeneous data. When learning on non-independent and identically distributed (non I.I.D.) data, due to the phenomenon of client drift, the accuracy of the global model of federated learning will be significantly reduced. 2. Lack of personalized solutions. In the original federated learning setting, a single global shared model is trained to adapt to the "averaged clients". In the case where there are obvious differences in the data distributions of each client, it is difficult for a single global model to handle the local distribution that is very different from the global distribution. Therefore, we need to perform personalized processing for each client to solve these two challenges, which are exactly the problems that personalized federated learning attempts to solve.
[0004] Personalized federated learning aims to meet the personalized needs of different clients, but most existing methods still fine-tune the client local models based on traditional federated learning. In this case, the server still aggregates the client models by parameter averaging, making the client models still need to adopt the same structure. However, in actual application scenarios, there are often significant differences in the hardware devices, storage capacities, and resource conditions of clients, and the limitation of model isomorphism is difficult to meet these diverse needs. Therefore, in order to break this constraint of model isomorphism, it is particularly important to design a more flexible personalized federated learning algorithm. Knowledge distillation, as a model-agnostic knowledge transfer method, exactly has such advantages. The teacher model guides the training of the student model by outputting soft labels without relying on the same model structure. Summary of the Invention
[0005] The present invention proposes a personalized federated learning method based on multi-teacher attention integrated distillation. This algorithm breaks the limitation of traditional federated learning on model isomorphism. By using integrated distillation and contrast learning techniques on the server side, clients can freely choose models suitable for themselves. The core idea of the present invention is that the server saves model copies of each client, and the global model uses the model outputs of clients with high similarity and strong relevance to itself through the attention mechanism for distillation learning. Finally, using the idea of contrast learning, the global model is compared with the models of each client, so that the client's model can learn the knowledge of the global model.
[0006] The process of training the personalized federated learning method based on multi-teacher attention integrated distillation includes the following steps:
[0007] Step S1: The server initializes the global model Θ s , allocates independent storage space for each client and saves its set of model copies and distributes the initial model parameters to all clients; the server maintains an unlabeled common dataset D i with the same feature space as the client's private dataset D pub ;
[0008] Step S2: In the t-th communication round where t ∈ {1, …, T}, the server randomly selects a client subset S according to a preset participation rate ρ ∈ (0.2, 1.0) t , and distributes the current global model to the clients in S after parameter compression encoding; t
[0009] Step S3: After receiving the global model, client i ∈ S t performs training for E ≥ 5 rounds on the local dataset D i , uses gradient clipping with a threshold of 1.0 and 8-bit parameter quantization to generate an updated local model Θ i , and encrypts and uploads the model parameters to the server;
[0010] Step S4: The server performs multi-teacher attention integrated distillation, calculates the similarity weights between the client models and the global model, generates consensus soft labels, and updates the global model;
[0011] Step S5: The server aligns the representation spaces of the clients and the global model through contrast learning to update the personalized model;
[0012] Step S6: Distribute the updated model to the clients and iterate the training.
[0013] Preferably, the implementation of multi-teacher attention integrated distillation in step S4 includes:
[0014] S41: When the server receives the locally trained model from the client, it saves it in the local model set of the server. Using the public dataset D pub Input the global model Θ s and each client model Generate the corresponding soft label distribution q(Θ s , x) and
[0015] S42: Calculate the cosine similarity between the client model and the global model based on the following formula to generate the normalized attention weights, calculate the similarity between the global model and each client model, and construct a similarity matrix
[0016]
[0017] S43: Aggregate the client soft labels weighted by the weights to construct a consensus soft label:
[0018]
[0019] S44: Minimize the KL divergence loss through the following formula Update the global model parameters:
[0020]
[0021] Furthermore, the S42 attention weight calculation module realizes dynamic knowledge selection through a i,s . In the formula, σ is a scaling factor used to adjust the sensitivity of the cosine similarity; cos(q(Θ i ), q(Θ s )) represents the cosine similarity between the soft labels of client i and the global model Θ s , and the cosine similarity metric is used to measure the output distribution similarity on the public dataset D pub ; a i,s is actually a Softmax normalization function that converts the cosine similarity into a probability distribution, and a i,s represents the contribution degree of each client to the global model, satisfying ∑ i∈[N] a i,s = 1.
[0022] Furthermore, in S44, each client model is used as an independent teacher model, and the global model is used as a student model; the KL divergence is used to measure the distribution difference between the consensus soft label and the output of the global model; the KL divergence loss is optimized using the gradient descent algorithm, and the learning rate η ∈ (0.001, 0.1), and the update rule is and the value of the scaling factor σ is dynamically adjusted according to the dataset complexity.
[0023] Preferably, in step S5, contrastive learning alignment is performed. Specifically, effective feature representations are learned by comparing the similarities and differences between data samples, enabling the model to learn to distinguish between similar samples (positive sample pairs) and dissimilar samples (negative sample pairs). The implementation steps include:
[0024] S51: Extract the intermediate representation r of the model for the data batch of the common dataset t (teacher model) and r s (student model), with the goal of maximizing the similarity of positive sample pairs while minimizing the similarity of negative sample pairs;
[0025] S52: Calculate the cosine similarity contrast loss
[0026]
[0027] where sim(·) represents the similarity function, and in this invention, the cosine similarity is adopted. That is, sim(r t , r s ) = r t ·r s / ‖r t ‖‖r s ‖. The temperature parameter τ is used to control the smoothness of the loss function; on the server side, the global model and the client model generate the representations r t and r s from the data batch given by the common dataset. The positive sample is the similarity between the representation r t of the global model and the representation r s of the client model, and the negative sample is the representation of other samples in the same batch;
[0028] S53: Use the gradient descent method to update the client model parameters and align its feature representation space with the global model:
[0029]
[0030] where η is the learning rate. After the client model is aligned with the global model, the updated personalized model is sent back to client i again. BRIEF DESCRIPTION OF THE DRAWINGS
[0031] Figure 1 is the system architecture diagram of the present invention;
[0032] Figure 2 is the overall flowchart of the present invention. DETAILED DESCRIPTION OF THE EMBODIMENTS
[0033] Next, the technical solutions in the embodiments of the present invention will be clearly and completely described in conjunction with the accompanying drawings in the embodiments of the present invention. Obviously, the described embodiments are only a part of the embodiments of the present invention, rather than all the embodiments. All other embodiments obtained by those of ordinary skill in the art based on the embodiments of the present invention without creative efforts shall fall within the protection scope of the present invention.
[0034] A personalized federated learning method based on multi-teacher attention integrated distillation. A typical system architecture is as Figure 1 shown, and the flowchart of the embodiment is as Figure 2 shown, including the following steps:
[0035] Step S1: The server initializes the global model Θ s , allocates independent storage space for each client and saves its set of model copies and distributes the initial model parameters to all clients; the server maintains an unlabeled common dataset D i consistent with the feature space of the client's private dataset D pub ;
[0036] Step S2: In the t-th communication round where t ∈ 1,..., T, the server randomly selects a subset S of clients according to a preset participation rate ρ ∈ (0.2, 1.0) t , and distributes the current global model after parameter compression encoding to the clients in S t ;
[0037] Step S3: After receiving the global model, client i ∈ S t performs training for E ≥ 5 rounds on the local dataset D i , uses gradient clipping with a threshold of 1.0 and 8-bit parameter quantization to generate the updated local model Θ i , and encrypts and uploads the model parameters to the server;
[0038] Step S4: The server performs multi-teacher attention integrated distillation, calculates the similarity weights between the client models and the global model, generates consensus soft labels, and updates the global model;
[0039] Specifically, the implementation of multi-teacher attention integrated distillation includes four steps:
[0040] S41: When the server receives the locally trained model of the client, it saves it in the server's local model set. Using the common dataset D pub as the input to the global model Θ s and each client model to generate the corresponding soft label distributions q(Θ s , x) and
[0041] S42: Adjust the similarity sensitivity through a scaling factor σ ∈ [0.1, 1], calculate the cosine similarity between the client model and the global model based on the following formula, generate normalized attention weights, calculate the similarity between the global model and each client model, and construct a similarity matrix
[0042]
[0043] where σ is a scaling factor used to adjust the sensitivity of the cosine similarity; cos(q(Θ i ), q(Θ s )) represents the cosine similarity between client i and the soft label of the global model Θ s , and uses the cosine similarity metric to measure the similarity of the output distributions on the common dataset D pub ; a i,s is actually a Softmax normalization function that converts the cosine similarity into a probability distribution, and a i,s represents the contribution degree of each client to the global model, satisfying ∑ i∈[N] a i,s = 1.
[0044] S43: Aggregate the client soft labels weighted based on the weights to construct a consensus soft label:
[0045]
[0046] S44: Use each client model as an independent teacher model and the global model as a student model; use the KL divergence to measure the distribution difference between the consensus soft label and the output of the global model; minimize the KL divergence loss through the following formula to update the global model parameters:
[0047]
[0048] where the KL divergence loss optimization uses the gradient descent algorithm, the learning rate η ∈ (0.001, 0.1), and the update rule is and the value of the scaling factor σ is dynamically adjusted according to the complexity of the dataset.
[0049] Step S5: The server aligns the representation spaces of the client and the global model through contrastive learning and updates the personalized model;
[0050] Specifically, the server learns effective feature representations by comparing the similarities and differences between data samples, enabling the model to learn to distinguish between similar samples (positive sample pairs) and dissimilar samples (negative sample pairs). The implementation steps include:
[0051] S51: Extract the model intermediate representation r of the data batch of the common data set t (teacher model) and r s (student model), with the goal of maximizing the similarity of positive sample pairs while minimizing the similarity of negative sample pairs;
[0052] S52: Calculate the cosine similarity contrast loss
[0053]
[0054] Among them, sim(·) represents the similarity function, and the present invention adopts the cosine similarity. That is, sim(r t , r s ) = r t ·r s / ‖r t ‖‖r s ‖. The temperature parameter τ is used to control the smoothness of the loss function; on the server side, the global model and the client model generate the representations r t and r s according to the data batch given by the common data set. The positive sample is the similarity between the representation r t of the global model and the representation r s of the client model, and the negative sample is the representation of other samples in the same batch;
[0055] S53: Use the gradient descent method to update the client model parameters and align its feature representation space with the global model:
[0056]
[0057] Among them, η is the learning rate. After the client model is aligned with the global model, the updated personalized model is sent back to client i again.
[0058] Step S6: Send the updated model to the client and perform iterative training.
[0059] Although the embodiments of the present invention have been shown and described, for those of ordinary skill in the art, it can be understood that various changes, modifications, substitutions, and variations can be made to these embodiments without departing from the principles and spirit of the present invention. The scope of the present invention is defined by the appended claims and their equivalents.
Claims
1. A personalized federated learning method based on multi-teacher attention integrated distillation, characterized in that , including the following steps: S1: The server initializes the global model Θ s and saves copies of each client model and distributes the initial model to the clients; S2: At communication round t ∈ {1,..., T}, the server samples the client set S t , and sends the current global model S3: Client \(i\in S\) t Uses the local dataset \(D\) i To train the local model \(\Theta\) i , and uploads the updated model parameters to the server; S4: The server performs multi-teacher attention integrated distillation, calculates the similarity weights between the client model and the global model, generates consensus soft labels, and updates the global model; S5: The server updates the personalized model by aligning the representation spaces of the client and the global model through contrastive learning; S6: The updated model is sent to the client for iterative training.
2. The method according to claim 1, characterized in that, The implementation of multi-teacher attention integrated distillation in step S4 includes: S41: Utilize the public dataset D pub Input the global model Θ s and each client model Generate the corresponding soft label distribution q(Θ s , x) and S42: Calculate the cosine similarity between the client model and the global model based on the following formula to generate normalized attention weights: S43: Weightedly aggregate the client soft labels based on the weights to construct consensus soft labels: S44: Minimize the KL divergence loss through the following formula Update the global model parameters: 。 3. The method according to claim 2, wherein The attention weight calculation module in step S42 realizes dynamic knowledge selection through a i,s In the formula, σ is a scaling factor used to adjust the sensitivity of the cosine similarity; cos(q(Θ i ),q(Θ s )) represents the cosine similarity between the soft labels of client i and the global model Θ s , and uses the cosine similarity metric to measure the similarity of the output distributions on the common dataset D pub ; a i,s is actually a Softmax normalization function that converts the cosine similarity into a probability distribution, and a i,s represents the contribution degree of each client to the global model, satisfying ∑ i∈[N] a i,s = 1.
4. The method according to claim 2, wherein In step S44, each client model is used as an independent teacher model, and the global model is used as a student model; the KL divergence is used to measure the distribution difference between the consensus soft label and the output of the global model; the KL divergence loss optimization uses the gradient descent algorithm, and the learning rate η ∈ (0.001, 0.1), and the update rule is Moreover, the value of the scaling factor σ is dynamically adjusted according to the dataset complexity.
5. The method according to claim 1, wherein For the contrastive learning alignment in step S5, specifically, effective feature representations are learned by comparing the similarities and differences between data samples, enabling the model to learn to distinguish similar samples (positive sample pairs) and dissimilar samples (negative sample pairs). The implementation steps include: S51: Extract the model intermediate representation r of the data batch of the common data set t (teacher model) and r s (student model), with the goal of maximizing the similarity of positive sample pairs and minimizing the similarity of negative sample pairs; S52: Calculate the cosine similarity contrastive loss Among them, sim(·) represents the similarity function, and the cosine similarity is adopted in the present invention. That is, sim(r t , r s ) = r t ·r s / ‖r t ‖‖r s ‖. The temperature parameter τ is used to control the smoothness of the loss function; on the server side, the global model and the client model generate representations r t and r s according to the data batches given by the common data set. The positive sample is the similarity between the representation r t of the global model and the representation r s of the client model, and the negative sample is the representation of other samples in the same batch; S53: Use the gradient descent method to update the client model parameters and align its feature representation space with the global model: Among them, η is the learning rate. After the client model is aligned with the global model, the updated personalized model will be sent back to client i again.