Federated Learning Method Based on Representation-Driven Head Clustering under Edge Devices

By adopting the characterization-driven head clustering method in federated learning, the decoupling model is used to generate data representations and perform weighted averages using the characterization generation module, which solves the problems of large communication overhead on edge devices and slow model convergence, and achieves efficient adaptability and accuracy of the model in a dynamic environment.

CN119312945BActive Publication Date: 2025-07-18HUBEI UNIV
View PDF 2 Cites 0 Cited by

Patent Information

Application Number
CN202411438940.9
Authority / Receiving Office
CN · China
Patent Type
Patents(China)
Current Assignee / Owner
Filing Date
2024-10-15
Publication Date
2025-07-18
Estimated Expiration
2044-10-15

AI Technical Summary

Technical Problem

The existing federated learning methods have problems such as large communication overhead, slow model convergence speed and cluster inadequacy on edge devices, and are difficult to effectively solve in resource-constrained environments.

Method used

The head clustering method based on characterization drive is adopted to decouple the model into a shared header and a characterization generation module. The characterization generation module generates data characterization and weighted average, and the client cluster is used to use characterization distance. The central server trains the shared header in the cluster and dynamically adjusts it to reduce traffic.

Benefits of technology

While ensuring the accuracy of the model, it significantly reduces communication overhead, enhances the adaptability of the model in a dynamic environment and the applicability of resource-constrained devices, and improves the generalization ability of the model.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN119312945B_ABST
    Figure CN119312945B_ABST
Patent Text Reader

Abstract

The present invention discloses a federated learning method based on representation-driven head clustering for edge devices. Core features are extracted from local raw data through representation learning, and only the data representations are uploaded to the server, which can reduce communication overhead while retaining the core information of the original data and ensuring the accuracy of the clustering process. The server trains the shared layer within the clusters based on these representations, thereby ensuring the accuracy of the model. The present invention decouples the model into a representation generation module and a shared head. The representation generation module is used to generate data representations; the shared head improves the generalization ability of the model by learning the knowledge of different clients. The central server dynamically clusters the clients according to the representations uploaded in each round, and can adjust the clustering structure in a timely manner according to the changes in the data in each round of iteration to ensure the accuracy and adaptability of the clustering. Finally, the central server uses the data representations of the clients within the clusters to train the shared head of each cluster and sends the trained shared head to the clients to enhance the generalization ability of the model.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the technologies of federated learning and artificial intelligence, and particularly to a federated learning method based on representation-driven head clustering under edge devices. Background Art

[0002] Mobile edge devices, such as smartphones, tablets, Internet of Things devices, and drones, generate a large amount of data every day, but due to privacy protection restrictions, it cannot be fully utilized. Federated learning is a distributed machine learning method to solve this problem. It allows devices to jointly train a model without sharing the original data, thus effectively protecting data privacy. However, traditional federated learning algorithms (such as FedAvg) have a significant problem of slow model convergence speed when facing data heterogeneity. In addition, some personalized federated learning methods adjust the model to adapt to local data, although this alleviates the problem, but still faces the challenges of local model overfitting and slow model convergence. To address this challenge, some studies have proposed personalized federated learning, whose core purpose is to train a model suitable for local data for each client while maintaining data privacy and security. The methods of personalized federated learning include model layering, transfer learning, meta-learning methods, and knowledge distillation, etc. However, most of these methods rely on each client independently using its own small and specific local dataset to adjust the model, which limits the ability to learn from the data of other clients, and thus leads to local model overfitting and slow model convergence. Clustering federated learning is proposed as a new strategy, which realizes cross-client knowledge sharing by clustering clients with similar data distributions. However, most of the existing CFL methods rely on initial clustering, which requires a large amount of communication and computation before the training starts, and after the clustering is completed, the clients are fixed in their respective clusters, which limits the adaptability of the system in a dynamic environment.

[0003] Although recent studies, IFCA and AICFL, have achieved dynamic adjustment of clustering during the federated learning training process. However, the server needs to broadcast the cluster model in each round of iteration, and the client needs to calculate the clustering identity locally, which increases the communication cost, especially on resource-constrained edge devices, resulting in a decrease in model training efficiency. Another study, HiCFL, performs client clustering by calculating the cosine similarity of the fully connected layer or convolutional layer of the model, introducing the concept of model stability. Therefore, it does not need to communicate the entire model between the server and the client, and reduces the number of communication rounds, thus reducing the communication and computation overhead. However, after the clustering stage, the clustering result is not adjusted, and the entire model needs to be communicated for traditional FL training of each cluster, resulting in a large communication overhead.

[0004] For example, Patent CN114936595A discloses a model fine-tuning and head aggregation method in federated learning. The client needs to send the entire model (local representation + head) to the server, resulting in a large communication overhead. In this technical solution, the server globally aggregates the local representations, greatly reducing the personalization degree of the model. Another example is Patent CN118470412A, which discloses a federated object detection learning method based on representation enhancement and weighted aggregation in a cloud-edge-end environment. Uploading the entire client model to the server will significantly increase the communication overhead. At the same time, its representation adds an imbalance factor (for the scenario of few-shot learning, there is a class imbalance in few-shot samples), and some strengthening operations need to be performed on the imbalance factor. In the final model aggregation, this will expose the information of the original data.

[0005] Therefore, although existing technologies and methods have proposed strategies to improve communication overhead and dynamic clustering, there are still two major deficiencies: (1) how to reduce the communication volume while ensuring that the model accuracy is not affected; (2) how to ensure the accuracy and adaptability of clustering under limited data.

[0006] These problems have not been fully solved and have become key challenges in the existing technology. Summary of the Invention

[0007] Object of the Invention: The object of the present invention is to solve the deficiencies existing in the existing technology and provide a federated learning method based on representation-driven head clustering under edge devices.

[0008] Technical Solution: A federated learning method based on representation-driven head clustering under edge devices according to the present invention includes the following steps:

[0009] Step 1, generate data representations, that is, use a representation generation module on the edge device to generate representations of data samples; and perform weighted averaging on the data representations to obtain data prototypes. The specific method is as follows:

[0010] Step 1.1, after each round of local training on the edge device, decouple the local model into a shared head and a representation generation module

[0011]

[0012] In the above formula, ○ is a concatenation symbol, represents the local model after local training is completed, and respectively represent the local shared head and the representation generation module after decoupling the local model of the kth edge device in the tth round of federated learning;

[0013] Step 1.2. The edge device uses the local characterization generation module to generate the characterization of the data samples participating in this round of training

[0014]

[0015] In the above formula, x i represents the i-th data sample, represents the function of using the characterization generation module to generate the data sample x i ;

[0016] Step 1.3. Weight and average the characterizations of samples with the same label to obtain the data prototype (that is, the content generated by calculating the local data through the characterization generation module);

[0017]

[0018] In the above formula, represents the total number of sample data with label q participating in this round of training;

[0019] Step 2. Cluster the clients, that is, first calculate the characterization distance between the clients, and then group the clients with similar data distributions into one cluster;

[0020] Step 3. Based on the clients after clustering in Step 2, train the shared head in each cluster. For the client k belonging to multiple clusters, aggregate the shared heads of all its clusters, and then send them to the corresponding client (edge device). Then, use the FedAvg algorithm to summarize the heads of all clusters.

[0021] Furthermore, Steps 2 and 3 are executed on the central server. The performance of the central server is much stronger than that of the edge device, so that the edge device can participate in the calculation as little as possible, which helps to improve the overall speed.

[0022] Furthermore, the detailed method of Step 2 is as follows:

[0023] Calculate the characterization distance dis between the data prototypes received from the clients to obtain the similarity between the data distributions of the two clients;

[0024]

[0025] In the above formula, dis q,r represents the characterization distance between the q-label data prototype of the k-th client and the l-label data prototype of the r-th client, and ε represents the threshold hyperparameter; when the characterization distance is less than ε, these two data prototypes are clustered into the same cluster in this round. The total number of clusters after clustering is M, and the number of clients in the m-th cluster is |Cm.

[0026] Further, the detailed method of step 3 is as follows:

[0027] After the central server completes the clustering of all representations, use the in-cluster representations and the corresponding label q to train the globally shared head within the cluster

[0028]

[0029] In the above formula, η θ represents the global learning rate, R m represents the in-cluster representation, and y represents the label corresponding to the in-cluster representation;

[0030] Since an edge device k may belong to different clusters, then aggregate the shared heads of different clusters and take the weighted average to obtain a new round of shared head for edge device k;

[0031]

[0032] In the above formula, represents the shared layer after aggregation and update of client k, and M k is the set of all clusters to which client k belongs, represents the proportion of the samples of client k in the m-th cluster in the total samples, is the number of data samples of client k belonging to cluster m;

[0033] After the shared head of each cluster is trained, the central server aggregates the shared layers of each cluster to obtain the shared layer of the global model The formula is as follows:

[0034]

[0035] In the above formula, represents the proportion of the number of clients in the m-th cluster in the total number of clients, |Cm| is the number of clients in cluster m, and M is the total number of clusters.

[0036] Beneficial effects: Compared with the prior art, the present invention has the following advantages:

[0037] (1). The present invention proposes a representation-driven head clustering specifically designed for federated learning under edge devices. This algorithm allows the server to cluster edge devices using limited information during training, enhances the model's ability to adapt to the continuously changing data distribution of edge devices, reduces the server's computational overhead, and ensures the accuracy of the clustering results.

[0038] (2) In the present invention, representations are utilized to train the globally shared head in each cluster, eliminating the need to transmit model weights. This method reduces the amount of data transmitted while maintaining model accuracy, thereby enhancing the applicability of federated learning in resource-constrained mobile edge environments.

[0039] (3) The present invention can reduce communication overhead while ensuring model accuracy. Additionally, the present invention is not limited to model types and can be integrated into the federated learning framework without imposing an additional burden on edge devices. BRIEF DESCRIPTION OF THE DRAWINGS

[0040] Figure 1 Schematic diagram of the overall framework of the present invention;

[0041] Figure 2 Schematic diagram of the overall process of the present invention;

[0042] Figure 3 Schematic diagram of the distance between client prototypes in the 1st round of the embodiment;

[0043] Figure 4 Schematic diagram of the distance between client prototypes in the 50th round of the embodiment;

[0044] Figure 5 Schematic diagram of the distance between client prototypes in the 300th round of the embodiment

[0045] Figure 6 Schematic diagram of the accuracy comparison of the embodiment in dynamic and static scenarios of CIFAR-10;

[0046] Figure 7 Schematic diagram of the accuracy comparison of the embodiment in dynamic and static scenarios of CIFAR-100;

[0047] Figure 8 Schematic diagram of the analysis of different hyperparameter thresholds ε for the embodiment on CIFAR-10;

[0048] Figure 9 Schematic diagram of the analysis of different hyperparameter thresholds ε for the embodiment on CIFAR-100;

[0049] Figure 10 Schematic diagram of the change of model accuracy with the threshold ε for the embodiment on CIFAR-10;

[0050] Figure 11 Schematic diagram of the change of model accuracy with ε for the embodiment on CIFAR-100. DETAILED DESCRIPTION OF THE INVENTION

[0051] The technical solution of the present invention will be described in detail below, but the protection scope of the present invention is not limited to the described embodiments.

[0052] As Figure 1 and Figure 2 shown, the federated learning method based on representation-driven head clustering under edge devices of the present invention is characterized by comprising the following steps:

[0053] Step 1, generating data representations, that is, generating representations of data samples by using a representation generation module on edge devices; and performing weighted averaging on the data representations to obtain data prototypes. The specific method is as follows:

[0054] Step 1.1, after each round of local training is completed on an edge device, decouple the local model into a shared head and a representation generation module

[0055]

[0056] In the above formula, ○ is a concatenation symbol, represents the local model after local training is completed, and respectively represent the local shared head and the representation generation module after decoupling of the local model of the kth edge device in the tth round of federated learning;

[0057] Step 1.2, the edge device uses the local representation generation module to generate representations of data samples participating in this round of training

[0058]

[0059] In the above formula, x i represents the ith data sample, represents the function of using the representation generation module to generate the data sample x i ;

[0060] Step 1.3, perform weighted averaging on the representations of samples with the same label to obtain a data prototype

[0061]

[0062] In the above formula, represents the total number of sample data with label q participating in this round of training;

[0063] Step 2, clustering the clients, that is, first calculating the representation distances between the clients, and then grouping the clients with similar data distributions into one cluster;

[0064] Step 3: Based on the clients clustered in Step 2, train the shared head in each cluster. For the client k that belongs to multiple clusters, aggregate the shared heads of all its clusters, then send them to the corresponding client, and then use the FedAvg algorithm to summarize the heads of all clusters.

[0065] The detailed processes of Step 2 and Step 3 in this embodiment are both executed on the central server. And the central server is based on prototype clustering, uses prototypes within the cluster to train the shared head, and sends the head within the cluster to the corresponding client.

[0066] The detailed method of Step 2 in this embodiment is as follows:

[0067] Calculate the representation distance dis between the data prototypes received from the clients to obtain the similarity between the data distributions of two clients;

[0068]

[0069] In the above formula, dis q,r represents the representation distance between the q-label data prototype of the k-th client and the l-label data prototype of the r-th client, and ε represents the threshold hyperparameter; when the representation distance is less than ε, these two data prototypes can be clustered into the same cluster in this round.

[0070] The detailed method of Step 3 in this embodiment is as follows;

[0071] After the central server finishes clustering all the representations, use the in-cluster representations and the corresponding label q to train the global shared head within the cluster

[0072]

[0073] In the above formula, η θ represents the global learning rate, R m represents the in-cluster representation, and y represents the label corresponding to the in-cluster representation;

[0074] Then aggregate and weighted average the shared heads of different clusters to obtain a new round of shared head for the edge device k;

[0075]

[0076] In the above formula, represents the shared layer after aggregation and update for client k, and M k is the set of all clusters to which client k belongs, represents the proportion of the samples of client k in the m-th cluster to the total samples, is the number of data samples of client k belonging to cluster m;

[0077] After the shared head of each cluster is trained, the central server aggregates the shared layers of each cluster, and the formula is as follows:

[0078]

[0079] In the above formula, represents the proportion of the number of clients in the m-th cluster to the total number of clients. |Cm| is the number of clients in cluster m, and M is the total number of clusters.

[0080] Embodiment 1

[0081] In this embodiment, the relevant steps in the entire technical solution are algorithmized, which is shown as follows:

[0082] 1. Define and initialize parameters: Set the total number of clients as N, the number of edge devices in each round of communication as K, determine the total number of communication rounds as T, and the local learning rate as η l , and the global learning rate as η θ . For the t = 0 communication round, perform the following steps: For each edge device k, initialize the local model Decouple the parameters into the local shared head and the representation generation module Use the representation generation module to generate data representations Average the data representations to generate prototypes and upload them to the server; The server calculates the correlation dis between data representations using the Euclidean distance; Based on dis, cluster the edge devices, and assign edge devices with highly correlated data distributions to the same cluster M.

[0083] 2. Intra-cluster model shared head training: Use the intra-cluster representation R m and the corresponding label y on each cluster m to train the intra-cluster shared head For each communication round from t = 1 to T, perform the following steps:

[0084] Operations at the edge device side are as follows: For each selected client K, perform the following operations: Concatenate the downloaded shared head with the local representation generation module to update the local model, and use local data to perform gradient descent update on the local model to obtain new model parameters The edge device uses the new representation generation module to generate data representations, and after averaging them, uploads them to the server;

[0085] Operations at the server side are as follows: The server re-clusters the clients based on the new representations and retrains the shared head in each cluster For the edge device k belonging to multiple clusters, aggregate the shared headers of the clusters it belongs to to obtain the shared header of the edge device k and send it to the corresponding client.

[0086] 3. Global model aggregation: The server aggregates the shared headers of each cluster to obtain the updated global shared header

[0087] Embodiment 2

[0088] To verify the technical effect of the present invention, this embodiment conducts experiments on two publicly available real datasets, which are widely used to evaluate the performance of federated learning models: CIFAR-10 and CIFAR-100. All experiments are developed using Python 3.11 and PyTorch 2.2.2 and executed on a standard computing platform equipped with an NVIDIA GeForce RTX 4090 GPU and 24GB of RAM.

[0089] To simulate non-independent and identically distributed data, this embodiment sets up two groups of experiments on the CIFAR-10 dataset for non-independent and identically distributed data simulation; Setting 1: Each client selects two different data classes, and the number of users is set to 20; Setting 2: Each client selects four different data classes, and the number of users is set to 50; In addition, the number of each class is equal, and all clients have the same number of data. In CIFAR-100, the number of clients is set to 10, each client selects ten different data classes, and each client has 3000 pictures, with every 300 pictures belonging to the same class.

[0090] For the selection of the model, this embodiment applies a 4-layer CNN model, including two convolutional layers and two fully connected layers, to perform classification tasks on the CIFAR-10 and CIFAR-100 datasets.

[0091] And the technical solution DRCFL of the present invention is compared with 8 current popular benchmark methods, including FedAvg, LG-FedAvg, FedProto, FedGH, FedPAC, IFCA, PACFL, and FedCAC, and the average accuracy and communication cost of the local models are reported. The entire optimization process uses the method of stochastic gradient descent. The experimental results show that the technical solution DRCFL of the present invention is superior to these methods in multiple aspects such as the reliability, stability, and communication efficiency of the system.

[0092] The details of the experimental results of this embodiment are as follows:

[0093] (1). Characterization similarity verification

[0094] The data representation obtained through the representation learning process is used as the basis for client clustering and partitioning in the technical solution DRCFL of the present invention. As Figures 3 to 5 shown, preliminary experiments using Euclidean distance to calculate the distances between data representations of different classes in the CIFAR-10 dataset indicate that the data representations can accurately reflect the similarity of the original data distribution.

[0095] (2) Model performance comparison

[0096] Accuracy evaluation: In this embodiment, the local models of each client were tested multiple times, and the average value of the accuracies of all clients was taken, recording multiple groups of average test accuracies. In terms of communication cost, the number of communication rounds required for the model to converge was recorded, and the communication cost per round was calculated (i.e., the amount of data uploaded and downloaded by each client multiplied by the number of participating clients), and then the total communication cost was obtained. The total communication cost was calculated by multiplying the communication cost per round by the number of communication rounds. This experiment first evaluated the model performance of different federated learning methods on the CIFAR-10 dataset. In the experimental configuration, the number of users participating in federated learning was set to 20 and 50 respectively, and the number of classes was 2 and 4 respectively. The experimental results are shown in Table 1.

[0097] Table 1

[0098]

[0099]

[0100] Table 1 compares the test accuracy %, the number of communication rounds required for convergence, and the total communication overhead of different methods on CIFAR-10. The symbol "—" indicates that the algorithm did not converge. From the above results, when the number of clients N = 20, the technical solution DRCFL of the present invention achieved a test accuracy of 88.70%, and the communication overhead was only 10.90 MB. In contrast, the single-round communication overhead of FedProto was only 56.1 KB because this method only needed to upload the local prototype and receive the global prototype from the server. However, FedProto required 257 rounds to converge, much higher than 115 rounds of the technical solution DRCFL of the present invention, and the total communication overhead was 14.1 MB, about 29.36% higher than the technical solution DRCFL of the present invention. The FedGH method requires the client to upload the prototype and download the global prediction head, and broadcast the global prediction head every round, resulting in its total communication overhead being five times that of DRCFL.

[0101] When N = 50, the technical solution DRCFL of the present invention still maintains the lowest communication overhead, only 95.1 MB, and achieves a test accuracy of 70.95%. In contrast, although FedPAC performs best in terms of test accuracy, reaching 78.96%, its communication cost is much higher than that of DRCFL. FedPAC needs to transmit representations and models in each round, resulting in a total communication cost more than 500 times that of DRCFL. Regardless of the change in the number of clients, the technical solution DRCFL of the present invention always maintains the best performance in terms of communication overhead and achieves a relatively high test accuracy.

[0102] Table 2

[0103]

[0104]

[0105] Comparison of the average test accuracy %, number of communication rounds required for convergence, and total communication overhead on CIFAR-100 in Table 2. To evaluate the performance of the algorithms on diverse datasets, tests were conducted on the CIFAR-100 dataset. The experimental results in Table 2 show that the technical solution DRCFL of the present invention achieves the highest test accuracy of 61.09%. In contrast, the test accuracy of FedCAC is 60.18%, but its total communication cost reaches 21.0 GB. FedGH performs better in terms of communication overhead, with a total communication cost of 1.2 GB and a test accuracy of 59.74%, but its performance is limited by a slower convergence speed and communication strategy. The total communication cost of the technical solution DRCFL of the present invention is 0.2 GB, which is approximately 83.33% lower than that of FedGH.

[0106] Although existing CFL methods can usually converge within 100 rounds under the same experimental conditions, their communication cost per round is significantly higher than that of the technical solution DRCFL of the present invention, about 500 times that of the technical solution DRCFL of the present invention. Although the convergence speed of these methods is relatively fast, their total communication cost is higher. For example, the test accuracy of PACFL is 57.20%, but its total communication cost is 53.6 GB, about 268 times that of the technical solution DRCFL of the present invention. It can be seen that the technical solution DRCFL of the present invention can significantly reduce the communication cost while ensuring a relatively high test accuracy and performs excellently on various datasets.

[0107] (3) Comparison of test accuracy in static and dynamic scenarios

[0108] To simulate the dynamic changes in data distribution, in this experiment, on the CIFAR-10 dataset, 5 categories were initially selected, and 10 clients were set for the first 50 rounds, with each client assigned 2 categories. At the 50th round, 10 new clients with different data distributions (non-IID, 2 / 10) were introduced. On the CIFAR-100 dataset, 50 categories were initially selected, and 10 clients were set for the first 100 rounds, with each client assigned 10 out of these 50 categories; at the 100th round, 10 new clients with different data distributions (non-IID, 10 / 100) were introduced.

[0109] This embodiment aims to simulate the integration of new clients and evaluate whether DRCFL can quickly adapt to dynamic data distributions. The experimental results are as Figure 6 , Figure 7 shown. In the dynamic scenario of the CIFAR-10 dataset, 10 new clients were introduced at the 50th round, resulting in a decrease in test accuracy. This decrease is attributed to the significant difference between the data distributions of the new clients and the original data distribution, and the model needs to be adjusted to adapt to these changes. After 25 rounds of adaptation, the test accuracy quickly recovered and gradually approached the level in the static scenario. A similar trend was also observed in the experiment on the CIFAR-100 dataset. These results indicate that the technical solution DRCFL of the present invention can effectively enable the model to adapt to changing data distribution scenarios.

[0110] (4), Hyperparameter Threshold Analysis

[0111] In this embodiment, cross-validation was performed to cross-validate the optimal range of hyperparameter thresholds on CIFAR-10 (non-IID, 2 / 10) and CIFAR-100 (non-IID, 10 / 100). Specifically, for the CIFAR-10 dataset, the thresholds were set to {3.0, 5.0, 5.5, 6.0, 6.5, 7.0, 16.0}; for the CIFAR-100 dataset, the thresholds were set to {5.0, 5.5, 6.0, 6.3, 6.5, 7.0, 9.0}. Figures 8 - 11 shows the impact of different thresholds on the model accuracy and convergence speed. The results show that the threshold has little impact on the final convergence accuracy, and the difference between the highest and lowest convergence accuracies is no more than 1%. However, the threshold has a significant impact on the convergence speed of the model. For example, when the threshold is set between 5.5 - 6.5, compared with the case where the threshold is 3.0, the number of rounds required for the model to converge is reduced by approximately 35% - 56%, and there are obvious fluctuations.

[0112] In summary, the present invention extracts core features from local raw data through representation learning and only uploads the data representations to the server, which can reduce communication overhead while retaining the core information of the raw data to ensure the accuracy of the clustering process. The server trains the shared layer within the cluster based on these representations, thereby ensuring the accuracy of the model. The present invention decouples the model into a representation generation module and a shared head. The representation generation module is used to generate data representations; the shared head improves the generalization ability of the model by learning the knowledge of different clients. The central server dynamically clusters the clients according to the representations uploaded in each round, and can timely adjust the clustering structure according to the changes in the data in each round of iteration to ensure the accuracy and adaptability of the clustering. Finally, the central server uses the data representations of the clients within the cluster to train the shared head of each cluster and sends the trained shared head to the clients to enhance the generalization ability of the model.

Claims

1. A federated learning method based on representation-driven head clustering under edge devices, characterized in that Including the following steps: Step 1: Generate data representations, that is, use a representation generation module on the edge device to generate representations of data samples; and perform weighted averaging on the data representations to obtain data prototypes. The specific method is as follows: Step 1.

1. After each round of local training is completed, the edge device decouples the local model into a local shared head and a representation generation module In the above formula, is the splicing symbol, represents the local model after local training is completed, and respectively represent the local shared head and the feature generation module after decoupling the local model of the k-th edge device in the t-th round of federated learning; Step 1.2: The edge device uses the local feature generation module to generate the features of the data samples participating in this round of training In the above formula, x i represents the i-th data sample, represents the function for generating the data sample x i using the feature generation module; Step 1.3: Weighted average the sample representations with the same label to obtain a data prototype In the above formula, represents the total number of sample data with label q participating in this round of training; Step 2: Cluster the clients, that is, first calculate the representation distance between the clients, and then group the clients with similar data distributions into one cluster; here, the representation distance refers to calculating the representation distance dis between the data prototypes received from the clients to obtain the similarity between the data distributions of two clients Step 3: Based on the clients clustered in Step 2, train the shared heads in each cluster. For client k that belongs to multiple clusters, aggregate the global shared heads of all its clusters, then send them to the corresponding client, and then use the FedAvg algorithm to summarize the heads of all clusters.

2. The federated learning method based on representation-driven head clustering under edge devices according to claim 1, characterized in that Execute Step 2 and Step 3 on the central server.

3. The federated learning method based on representation-driven head clustering under an edge device according to claim 1 or 2, characterized in that, The detailed method of Step 2 is as follows: Calculate the prototype of the data received from the client to obtain the representation distance dis therebetween, and obtain the similarity between the two client data distributions; In the above formula, dis q,r represents the representation distance between the q-label data prototype of the k-th client and the r-label data prototype of the l-th client, and ε represents the threshold hyperparameter; when the representation distance is less than ε, these two data prototypes are clustered into the same cluster in this round, the total number of clusters after clustering is M, and the number of clients in the m-th cluster is |Cm|.

4. The federated learning method based on representation-driven head clustering under edge devices according to claim 1, characterized in that, The detailed method of Step 3 is as follows; After the central server has completed all the representations of clustering, the global shared head within the cluster is trained using the in-cluster representations and the corresponding label q In the above formula, η θ represents the global learning rate, R m represents the intra-cluster representation, and y represents the label corresponding to the intra-cluster representation; Then aggregate and perform weighted averaging on the shared heads of different clusters to obtain a new round of shared heads for edge device k; In the above formula, represents the shared layer after the aggregated update of client k, and M k is the set of all clusters to which client k belongs, represents the proportion of the samples of client k in the m-th cluster to the total samples, is the number of data samples of client k belonging to cluster m; After the training of the shared head of each cluster is completed, the central server aggregates the shared layers of each cluster to obtain the shared layer of the global model The formula is as follows: In the above formula, represents the proportion of the number of clients in the m-th cluster to the total number of clients, |Cm| is the number of clients in cluster m, and M is the total number of clusters.

Citation Information

Patent Citations

  • Federal learning-based model training method, system and device, and storage medium

    CN117744832A

  • Personalized federal learning method based on comparative learning and conditional calculation

    CN118396082A