Self-adaptive clustering federated learning method based on prototype network in heterogeneous scene

By using the adaptive clustering method based on prototype networks in federated learning, dynamically detecting changes in client data distribution and adjusting clustering results, the problem of inaccurate clustering results in the existing technology is solved, and efficient personalized model training and new client integration are achieved.

CN119989034APending Publication Date: 2025-05-13BEIJING UNIV OF TECH
View PDF 0 Cites 3 Cited by

Patent Information

Application Number
CN202510067738.8
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-01-16
Publication Date
2025-05-13

AI Technical Summary

Technical Problem

The existing federated learning clustering method is difficult to dynamically detect changes in client data distribution, resulting in inaccurate clustering results, affecting the prediction performance of personalized models, and unable to effectively handle scenarios joined by new clients.

Method used

Adaptive clustering method based on prototype network is adopted, and the client generates a prototype representation in the prototype network space through the client, calculates the loss function to detect the data distribution changes, dynamically adjusts the clustering results, and calculates the probability of which cluster it belongs to for the new client, and decides which cluster it joins or forms its own cluster.

Benefits of technology

Dynamic clustering is realized, computing overhead is reduced, the prediction performance of personalized models is improved, and the scalability of federated learning is enhanced. It is suitable for small sample data scenarios.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN119989034A_ABST
    Figure CN119989034A_ABST
Patent Text Reader

Abstract

The invention discloses an adaptive clustering federated learning method based on a prototype network in a heterogeneous scene, and the method comprises the steps: obtaining the class prototype information of different classes in each client through information extraction, constructing the user prototype representation of the client, employing a cosine distance as the similarity measurement of different clients, and carrying out the iterative clustering based on the similarity measurement. The client evaluates whether data distribution changes or not after a certain round of training, samples new data of the client to obtain a query set, generates prototype representation of each category in the query set, calculates a loss function of the client in the query set, and if the loss function value is too large, the client quits current cluster training, and if the loss function value is too large, the client does not quit current cluster training. Adding federated learning training as a new federated learning participant; when a new client joins in federated learning training, the central server is responsible for calculating the probability that the new client belongs to each cluster and further determining which cluster the new client joins or forms a cluster by itself, so that the expandability of clustering federated learning is increased, and the prediction performance of a personalized model is improved.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention belongs to the field of distributed machine learning, and specifically relates to an adaptive clustering federated learning method based on a prototype network in a heterogeneous scenario. The method allows a central server to cluster clients according to their user prototype representations, and the client calculates the loss function of the user prototype representation on a newly added data set to detect changes in local data distribution, thereby achieving dynamic clustering. In addition, the scenario in which new clients join federated learning training is expanded to improve the personalized performance of federated learning. Background Art

[0002] Federated learning is a distributed machine learning technology that can be trained between multiple decentralized clients holding local data samples without exchanging private data. However, due to the differences in identity, behavior, environment, etc. between clients, their data is non-independent and identically distributed (Non-IID), resulting in large deviations in the performance of the global model on different clients. For some clients, the performance of the global model may be worse than that of the local model trained only on their private data, depriving them of the motivation to participate in federated learning. Personalized federated learning, as a subfield of federated learning, can solve the performance degradation of the global model caused by data heterogeneity and provide each client with a personalized model that is more adapted to its local data distribution. However, the data in a single client is extremely limited, and its personalized model may be biased or overfit.

[0003] Clustering federated learning is a compromise between local training and global training. It assigns clients with homogeneous data to the same cluster. The cluster model extracts more useful knowledge from the homogeneous data, alleviating the problem of insufficient local training data, thus having better prediction performance. At present, studies have proposed a hierarchical CFL scheme, which achieves clustering by separating clients through an optimal binary algorithm based on cosine similarity. Another study proposed an iterative federated clustering scheme IFCA, which uses a greedy strategy to periodically iterate cluster identity estimation. However, the above clustering methods all assume that the data distribution of the client does not change during the period of participating in the federated learning training, there are no new clients joining the federated learning training, and the number of clusters K needs to be estimated a priori. Responding to distribution changes only by re-clustering will bring huge computational overhead.

[0004] Clustering algorithms for federated learning face the following challenges: First, in actual applications of federated learning, the data distribution of clients is not static. The training data of each client gradually arrives in the form of "streaming data", and its data distribution also changes over time. Outdated data distribution will lead to inaccurate clustering results, thereby affecting the predictive performance of the personalized model of federated learning. Therefore, it is necessary to dynamically detect changes in data distribution at runtime. Second, most existing clustering methods for federated learning divide clusters based on the Euclidean distance between model parameters. Model parameters cannot directly represent the actual data distribution of the client and are difficult to apply to small sample data scenarios. Summary of the invention

[0005] In order to improve the performance of personalized federated learning models in heterogeneous scenarios and solve the above problems of the prior art, the present invention proposes an adaptive clustering federated learning method based on prototype networks.

[0006] The prototype network-based adaptive clustering federated learning method proposed in this invention mainly involves two types of participants: a central server and multiple clients participating in the training. The clients pre-train the prototype network and perform local model training. The central server is responsible for clustering the clients and aggregating the models uploaded by the clients in the cluster to obtain a personalized global model. The method is divided into three stages:

[0007] 1) Initialization clustering phase: local data set of client i The representation of internal samples is {(x1,y1),…(x j ,y j ),…,(x n ,y n )}, where x j is the vector representation of the sample, y j is the category label. For each category c, z sample points are selected from the total sample set to form the support set S c ,Furthermore, using information extraction Generate a prototype representation of the category:

[0008]

[0009] Where |S c | is the support set S c The size of It can be any information extraction method such as CNN, LSTM, BERT, etc. The class prototype maps the original samples of each category to the prototype network space. Choosing a suitable extraction method can make the samples of the same category closer in the prototype network space and the samples of different categories farther away. In this way, the user prototype representation of client i can be obtained C is the total number of categories of all client data samples. After the client obtains the user prototype representation, it sends it to the central server.

[0010] After receiving the user prototype representations of all clients, the central server continuously merges the cosine distances between the user prototype representations The two smallest clients generate a hierarchical clustering tree. After the clustering is completed, a cross-cut at any level can be made according to the needs to obtain the specified number of clusters K. Then, the central server sends the cluster number to each client and represents the average value of the user prototype of each cluster. As the cluster prototype representation of the cluster, m is the number of clients in the cluster.

[0011] 2) Adaptive dynamic clustering stage: The present invention considers a more practical problem. In the federated learning process, client data arrives dynamically over time. In the tth round of global training, the local data set of client i where d i (t) represents the data that arrives at the client between the t-1th round and the tth round of training. In particular, d i (0) represents the local data of client i before participating in federated learning training. After initializing clustering, the client evaluates whether the local data distribution has changed every σ rounds of training. Specifically, for each category c, client i selects z sample points from the sample set that arrives at the client between the t-σ round and the t round of training to form the query set Q c , similar to initializing clustering, generate a prototype representation of the category Calculate the number of clients i in the query set Q c The loss function is:

[0012]

[0013] Among them, c' is another category different from c, and the distance metric is When the loss function of client i is greater than the threshold δ, it indicates that the cluster prototype representation of the cluster In the query set Q c The prediction performance is poor, that is, the data distribution arriving at the client between the t-σth round and the tth round of training is different from the local data distribution. Client i exits the current cluster training and joins the federated learning training as a new federated learning participant, and then re-executes stage 1) to join a cluster that is more suitable for the local data distribution.

[0014] When a new client g joins the training, it is necessary to complete the local prototype network training according to stage 1) to obtain its user prototype representation U g , and send it to the central server, which is responsible for calculating the probability that the client g belongs to each cluster h:

[0015]

[0016] Where K is the number of clusters, and the cluster number h with the maximum probability and a value greater than or equal to the threshold γ is selected and sent to the client g. After that, the client g participates in the personalized training of cluster h. If the maximum probability is less than the threshold γ, it means that the data distribution of the current client g is not similar to that of all clusters. The client g will form a new cluster alone, and the central server will also represent its user prototype U g As a cluster prototype representation of this cluster, the number of clusters K=K+1.

[0017] 3) Aggregation phase: The central server sends the initialized personalized global model parameters to all clients in each cluster h After that, a clients are randomly selected in the tth round of global training of cluster h, and the selected client i trains the local model by gradient descent method

[0018]

[0019] Where η is the local model learning rate, express right Find the gradient, Represents the local loss function of client i:

[0020]

[0021] in represents the loss of the jth sample. The central server aggregates the local models to obtain the global personalized model of cluster h.

[0022]

[0023] The central server performs T rounds of aggregation on each cluster, and finally obtains the personalized models {θ1,…,θ K}.

[0024] Beneficial effects of the present invention:

[0025] 1. The present invention mainly solves the problem that outdated data distribution on the client leads to inaccurate clustering results, which affects the personalized performance of federated learning. The present invention uses information extraction to obtain class prototype information of different categories in each client, and then constructs the user prototype representation of the client, uses cosine distance as the similarity measure of different clients, and performs iterative clustering based on this. Every σ rounds of training evaluates whether the data distribution has changed, samples the query set from the new data of the client, and generates a prototype representation of each category in the query set, calculates the loss function of the client in the query set, and if the loss function value is too large, the client exits the current cluster training and joins the federated learning training as a new federated learning participant, thereby joining a cluster that is more suitable for the local data distribution, thereby achieving dynamic clustering without the need to re-cluster all clients, reducing computational overhead.

[0026] 2. The present invention proposes an efficient clustering mechanism when a new client joins federated learning training. The central server is responsible for calculating the probability that the new client belongs to each cluster, and then decides which cluster it joins or forms its own cluster, thereby increasing the scalability of cluster federated learning.

[0027] 3. The present invention takes into account the small sample problem in the actual application scenarios of federated learning (such as IoT devices), and implements clustering based on the prototype network rather than the model parameters. The model parameters cannot directly represent the actual data distribution of the client, and are affected by other client model parameters during federated learning training. The prototype network can reflect the actual data distribution of the client, and the clustering effect is better, thereby improving the prediction performance of the personalized model. BRIEF DESCRIPTION OF THE DRAWINGS

[0028] Figure 1 It is an interactive flow chart of the adaptive clustering federated learning method based on prototype network in heterogeneous scenarios;

[0029] Figure 2 Figure 2 is a diagram of an adaptive dynamic clustering method for federated learning based on prototype networks;

[0030] Figure 3 Schematic diagram of personalized federated learning aggregation method. DETAILED DESCRIPTION

[0031] The present invention will be further described below in conjunction with the accompanying drawings and specific embodiments.

[0032] The interactive flow chart of the adaptive clustering federated learning method based on prototype network in heterogeneous scenarios described in the present invention is as follows: Figure 1 As shown in the figure, there are mainly two types of participants involved: the central server and multiple clients participating in the training. The clients pre-train the prototype network and perform local model training. The central server is responsible for clustering the clients and aggregating the models uploaded by the clients in the cluster to obtain the personalized global model of all clusters. The specific implementation process includes the following steps:

[0033] Step 1: Initialize the clustering phase

[0034] 1) If Figure 2 As shown, client i selects z sample points from the total sample set for each category c to form a support set S c , using information extraction method Generate a prototype representation of the category

[0035] 2) Represented by each category prototype Get the user prototype representation U of client i i ;

[0036] 3) The client sends its user prototype representation to the central server;

[0037] 4) The central server receives the user prototype representations of all clients and continuously merges the cosine distance sim(U i ,U j ) The two smallest clients generate a hierarchical clustering tree;

[0038] 5) Make a cross-section at any level to obtain the specified number of clusters K;

[0039] 6) The central server returns the cluster number to each client and embeds the user of each cluster to represent the average value Serves as the cluster prototype representation of this cluster.

[0040] Step 2: Adaptive dynamic clustering stage

[0041] 1) If Figure 2 As shown, the client evaluates whether the local data distribution has changed every σ rounds of training. For each category c, z sample points are selected from the sample set that arrived at the client between the t-σ round and the t round of training of client i to form the query set Q c , generate a prototype representation of this category Client i computes the query set Q c The loss function on If the loss function If it is greater than the threshold δ, it indicates that the local data distribution has changed;

[0042] 2) If the local data distribution changes, client i exits the current cluster training and joins the federated learning training as a new federated learning participant, executing 3) in step 2;

[0043] 3) When a new client g joins the federated learning training, execute step 1 to complete the local prototype network training and obtain its user prototype representation U g , will U gThe central server calculates the probability p(h|g) that client g belongs to each cluster. The central server selects the cluster number h with the maximum probability and the value is greater than or equal to the threshold θ and sends it to client g. Client g participates in the personalized training of cluster h. If the maximum probability is less than the threshold θ, client g will form a new cluster alone, and the central server will represent its user prototype U g As a cluster prototype representation of this cluster, the number of clusters K=K+1.

[0044] Step 3: Aggregation phase

[0045] 1) If Figure 3 As shown, the central server sends the initialized personalized global model parameters to all clients in each cluster h.

[0046] 2) In the tth round of global training of cluster h, a client is randomly selected, and the selected client i trains the local model according to formula (4):

[0047] 3) The central server aggregates the local models through formula (6) to obtain the global personalized model of cluster h

[0048] 4) The central server repeats step 2) and 3) in step 3 until the global training rounds of each cluster reach the set value T, and finally obtains the personalized models {θ1,…,θ K}.

Claims

1. An adaptive clustering federated learning method based on prototype network in heterogeneous scenarios, characterized in that: The following steps are involved: Step 1: Initialize the clustering phase; 1) Client i selects z sample points from the total sample set for each category c to form a support set S c , using information extraction method Generate a prototype representation of the category 2) Represented by each category prototype Get the user prototype representation U of client i i ; 3) Client i sends its user prototype representation to the central server; 4) The central server receives the user prototype representations of all clients and continuously merges the cosine distance sim(U i , U j ) The two smallest clients generate a hierarchical clustering tree; 5) Make a cross-section at any level to obtain the specified number of clusters K; 6) The central server returns the cluster number to each client and embeds the user of each cluster to represent the average value As the cluster prototype representation of the cluster; Step 2: Adaptive dynamic clustering stage; 1) The client evaluates whether the local data distribution has changed every σ rounds of training; 2) If the local data distribution changes, client i exits the current cluster training and joins the federated learning training as a new participant, executing step 2, step 3); 3) Achieve efficient grouping when a new client g joins the federated learning training; Step 3: Aggregation stage; 1) The central server sends the initialized personalized global model parameters to all clients in each cluster h 2) Randomly select a clients in the tth round of global training of cluster h, and the selected client i trains the local model 3) The central server aggregates local models to obtain the global personalized model of cluster h 4) The central server repeats 2) and 3) in step 3 until the global training rounds of each cluster reach the set value T, and finally obtains the personalized models {θ1, ..., θ K }.

2. According to claim 1, the adaptive clustering federated learning method based on prototype network in heterogeneous scenarios is characterized in that: In step 1) 1) generate a prototype representation of the category The formula is as follows: where x j is the vector representation of the sample, y j is the category label, S c For each category c, z sample points are selected from the total sample set to form the support set. c | is the support set S c The size of It is any information extraction method, including CNN, LSTM, and BERT.

3. According to claim 1, the adaptive clustering federated learning method based on prototype network in heterogeneous scenarios is characterized in that: In step 2, 1) client i evaluates whether the local data distribution has changed, including the following steps: 1) For each category c, select z sample points from the sample set that arrives at the client i between the t-σth round and the tth round of training to form the query set Q c , generate a prototype representation of this category 2) Calculate the client i in the query set Q c The loss function If the loss function of client i is If it is greater than the threshold δ, client i exits the current cluster training.

4. According to claim 3, the adaptive clustering federated learning method based on prototype network in heterogeneous scenarios is characterized in that: 2) Calculate the client i in the query set Q c The loss function The formula is as follows: Among them, c' is a category different from c, For client i in the support set S c The prototype representation of category c in is: For client i in the query set Q c Prototype representation on category c, distance metric 5. According to claim 1, the adaptive clustering federated learning method based on prototype network in heterogeneous scenarios is characterized in that: In step 2, 3) the efficient grouping method when a new client g joins the federated learning training includes the following steps: 1) The new client g executes step 1 to complete the local prototype network training and obtains its user prototype representation U g ; 2) Will U g Send it to the central server, and the central server calculates the probability p(h|g) that client g belongs to each cluster; 3) The central server selects the cluster number h with the maximum probability and the value is greater than or equal to the threshold θ and sends it to the client g. The client g participates in the personalized training of cluster h. If the maximum probability is less than the threshold θ, the client g will form a new cluster alone, and the central server will represent its user prototype U g As a cluster prototype representation of this cluster, the number of clusters K=K+1.

6. The method for adaptive clustering federated learning based on prototype network in heterogeneous scenarios according to claim 5, characterized in that: 2) The central server calculates the probability p(h|g) that the client g belongs to each cluster using the following formula: Where K is the number of clusters, is the cluster prototype representation of cluster h, U g is the user prototype representation of client g.

Citation Information

Cited By

  • Personalized federal learning method and system based on dynamic hierarchical regulation and control

    CN120235219A

  • A personalized federated learning method and system based on dynamic hierarchical regulation

    CN120235219B

  • Federal semi-supervised domain adaptive time sequence learning method

    CN121561461A