Federal domain generalization method based on trainable prototype
By introducing adaptive marginal enhancement contrast learning and prototype diversity learning techniques in federated learning, the trainable prototype module is optimized, which solves the problem of performance degradation of federated learning methods on different data domains, and significantly improves the generalization ability and classification accuracy of the model.
Patent Information
- Application Number
- CN202510102247.2
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-01-22
- Publication Date
- 2025-05-23
AI Technical Summary
When facing different data domains, existing federated learning methods have problems such as insufficient inter-class distance, insufficient in-class consistency and insufficient robustness to unseen domains, resulting in a significant decline in the performance of the model on unknown data domains.
The federated domain generalization method based on trainable prototypes is adopted, and the trainable prototype module is optimized through adaptive marginal enhancement contrast learning and prototype diversity learning technology, improving the inter-class separation and intra-class consistency of the trainable prototype set, thereby enhancing the generalization ability of the model in unknown fields.
It significantly improves the generalization performance and classification accuracy of the model in different data distributions and complex changes, solves the problems of small inter-class distances and insufficient in-class consistency, and enhances the robustness of the model to unknown domains.
Smart Images

Figure CN120032203A_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the technical field of federated domain generalization methods, and in particular to a federated domain generalization method based on a trainable prototype. Background Art
[0002] In order to solve the dual challenges of data silos and privacy protection, federated learning has been proposed as a distributed machine learning technology. It allows multiple data owners to collaboratively train models without sharing original data and only exchange model parameters, thereby achieving knowledge sharing and privacy protection.
[0003] Traditional federated learning improves the performance of the global model by training local models locally on the client and aggregating model parameters on the server. Federated learning can converge to good performance under the assumption that the private data of all clients comes from the same domain. However, in many practical application scenarios, there are differences in the distribution of training data and test data, and the distribution of test data is unknown. When the private data of the client comes from different domains and the distribution of the test domain is unknown, the performance of the model on the unknown data domain will be significantly reduced, which is the "domain drift" phenomenon.
[0004] Existing federated learning schemes [Huang W, Ye M, Shi Z, et al. Rethinking federated learning with domain shift: A prototype view. In 2023 IEEE [C] / / CVF Conference on Computer Vision and Pattern Recognition (CVPR). 2023: 16312-16322.] proposed a prototype-based federated learning method, which captures domain diversity characteristics by constructing clustering prototypes and uses unbiased prototypes to provide stable and fair optimization objectives.
[0005] Huang et al. proposed a prototype-based federated learning method that has the following defects when facing the domain drift problem:
[0006] 1. Insufficient inter-class distance: Prototype-based federated learning methods capture the diversity of domains by clustering prototypes. However, the complexity and diversity of real data may cause data of different categories to be similar in certain feature dimensions. This makes the local prototypes of different classes relatively close in the feature space. When the clustering operation is performed based on data similarity, the prototype distance between classes is often close, which may lead to confusion and misclassification in classification tasks.
[0007] 2. Insufficient intra-class consistency: When the prototype-based federated learning method generates prototypes on the client, the data distribution of each client may have significant deviations, resulting in large differences between local prototypes of the same class. The clustered prototypes of the same class obtained by clustering also have large differences, resulting in insufficient intra-class prototype consistency. The unbiased prototype is obtained by averaging the clustered prototypes of the same class, but this method may not be able to fully capture and integrate the difference information between different clustered prototypes in the same class, affecting the classification accuracy of the model on the same class of data.
[0008] 3. Insufficient robustness to unseen domains: The prototype-based federated learning method aims to improve the performance of the model on clients with known data distribution, but does not consider the data distribution of potential unknown clients. Therefore, this method may perform poorly when facing unknown domains. Summary of the invention
[0009] The technical problem to be solved by the present invention is to provide a federated domain generalization method based on trainable prototypes in view of the deficiencies of the above-mentioned prior art, which is applied to the field of image recognition to complete corresponding classification tasks. In order to solve the problem of "domain drift", this method introduces domain generalization, aiming to train a model that can perform well in multiple fields with different data distributions under the premise of protecting privacy. Specifically, the method optimizes the trainable prototype module through adaptive marginal enhancement contrast learning and prototype diversity learning technology, improves the inter-class separability and intra-class consistency of the trainable prototype set, thereby enhancing the generalization ability of the model in unknown fields, and solving the problems of small inter-class distance and insufficient intra-class consistency in the existing methods.
[0010] In order to solve the above technical problems, the technical solution adopted by the present invention is: a federated domain generalization method based on a trainable prototype, involving data interaction between a server and several clients, including the following steps:
[0011] Step 1: Establish a trainable prototype module to generate and update the prototype representation of the category. The initialization parameters of the trainable prototype module include: the number of categories in the dataset, the number of prototypes for each category, the feature dimension, and the hidden layer dimension;
[0012] The trainable prototype module structure includes an embedding layer and a multi-layer nonlinear fully connected network;
[0013] The embedding layer is used to generate an initial embedding representation of the category feature prototype; according to the number of categories and the number of prototypes of each category in the model parameters, the combination of each category and the prototype within the class is mapped to a unique index, and the index is used as the input of the embedding layer to generate a corresponding embedding vector, and the dimension of each generated embedding vector is determined by the feature dimension specified in the model parameters; the embedding layer uses a normal distribution to initialize the embedding weights and defines the initial distribution of the embedding vectors, so that the embedding vectors have good initial separation in the feature space;
[0014] The multi-layer nonlinear fully connected network includes two hidden layers and one output layer; the output feature dimension of the hidden layer is determined by the hidden layer dimension specified in the model parameters, and the hidden layer performs nonlinear mapping on the embedded vector through two fully connected transformations and a ReLU activation function to gradually enhance the expressive power of the feature, and finally, the adjusted feature is mapped to the final prototype representation through the output layer;
[0015] Step 2: If the server and the client are interacting for the first time, the server sends the initialized global model parameters to the client. If the server and the client are not interacting for the first time, the server sends the global model parameters and the global prototype set to the client.
[0016] Step 3: If the client and the server are interacting for the first time, the client receives the initialized global model parameters sent by the server and initializes the local model. If the client and the server are not interacting for the first time, the client receives the global model parameters and the global prototype set sent by the server. The client uses its own source domain data and the received global prototype set to iteratively train the local model, and uploads the trained local model parameters and local prototype set to the server.
[0017] All clients share a model with the same structure, which includes two modules: feature extractor h and classifier f. The local training process of the client includes the following steps:
[0018] Step 3.1: During the local training process of client k, use the feature extractor h of the local model of client k k For the sample x in the client's own source domain data i Perform feature extraction and obtain the corresponding feature vector z i , the feature vector z i Input to the classifier f and calculate the classifier prediction probability distribution δ(f(z i )) and the corresponding true category y i The cross entropy loss function value of ;
[0019] Feature Extractor h k : The source domain data input by client k Extract feature maps to output space Among them, R V represents the input space of dimension V, is the source domain data of client k, is the source domain data of client k Mapped output space; get the corresponding eigenvector Among them, R D represents the output space with dimension D;
[0020] Classifier f: The feature vector z of client k i Mapped to M-dimensional vector space, generating an unnormalized prediction score f(z i ), by calculating the unnormalized prediction score f(z i ) Apply the softmax function to obtain the predicted probability distribution of each category δ(f(z i )) and calculate the classifier prediction probability distribution δ(f(z i )) and the corresponding true category label y i The cross entropy loss function value is as shown in the following formula:
[0021]
[0022] in, is the true category label y i The corresponding one-hot encoding vector has a dimension that matches the number of categories M. If y i =m, then The mth category of is taken as 1, and the rest are 0;
[0023] Step 3.2: Determine the global prototype set Is it empty? If the global prototype set If empty, set the global prototype contrast learning loss Is 0; if the global prototype set If not empty, the feature extractor h of the local model of client k is calculated. k The feature vector z of the extracted source domain sample i Global prototype comparison learning loss with the global prototype set By minimizing the global prototype contrastive learning loss Close the feature vector z of the source domain sample of client k i The similarity between the global prototype set of the same category and the feature vector z of the source domain sample of client k is pulled away i Similarity with the global prototype set of other categories, using the minimized global prototype contrastive learning loss Optimize the local feature extractor h of client k k , making the feature vector of the client k source domain samples closer to the local prototype of the correct category;
[0024] Global prototype contrastive learning loss As shown below:
[0025]
[0026] Among them, s(z i,a) is the feature vector z of client k i With the category prototype a∈g m The similarity measure of s(z i ,b) is the feature vector z of client k i With the category prototype b∈g m′ Similarity measure, g is the set of all category prototypes trained on the server side, g m is the prototype set belonging to category m, g m′ =gg m is the set of prototypes that do not belong to category m among all category prototypes;
[0027] Step 3.3: Compare the global prototype to the learning loss l GPCL and cross entropy loss The weighted combination is the total loss, minimizing the total loss and optimizing the client local model parameter θ k ;
[0028] Step 3.4: Loop through steps 3.1 to 3.3 until the total number of local training rounds on the client k is reached. k The output feature vector z i As the class prototype information, the average feature vector of all samples in category m is represented as the prototype of the mth class, and the local prototype set O is updated. k ;
[0029] Take the average of the feature vectors of the same category on client k to obtain the local prototype of category m on client k As shown in the following formula:
[0030]
[0031] in, is the dataset whose label belongs to category m in client k, is the j-th sample belonging to category m in client k, For sample x j The corresponding label, h k (x j ) is the local feature extractor h on client k k For sample x j The extracted feature vector, k = {1, 2, ..., K} is the client set, m = {1, 2, ..., M} is the category set, M is the number of categories, is the local prototype of category m on client k;
[0032] Then we get the local prototype set O of all categories of client k k , as shown in the following formula:
[0033]
[0034] Step 3.5: Upload client k’s local prototype set O k and local model parameters θ k Go to the server and wait for the next data interaction with the server;
[0035] Step 4: The server receives the local prototype sets and local model parameters uploaded by each client, performs weighted average of the local model parameters of each client according to the proportion of the local data volume of each client to the total data volume, and generates updated global model parameters; uses an unsupervised clustering algorithm to cluster the local prototype sets belonging to the same category to obtain clustered prototype sets of the same category; calculates the minimum Euclidean distance between the clustered prototype sets of each category, obtains the maximum inter-class distance between all categories, compares it with the pre-set threshold, and selects the smaller value of the two as the margin value for data interaction between the client and the server in this round;
[0036] The updated global model parameter θ′ is:
[0037]
[0038] Among them, N k is the number of samples on client k, and N is the total number of samples on all clients;
[0039] According to the clustering results, several representative prototypes are selected as clustering prototypes for each category, and the clustering prototype set C of all categories is obtained, as shown in the following formula:
[0040]
[0041] C={C 1 ,...,C m ,...C M}
[0042] Among them, C is the clustering prototype set of all categories, C m is the clustering prototype set of the mth class, For N m The representative prototype obtained by clustering the m-th local prototypes, J m is the number of cluster prototypes of the mth class, N m is the number of local prototypes of the mth type, Cluster is an unsupervised clustering algorithm operation, and the input m-type local prototype set O k Clustering is performed together to obtain J m Cluster centers;
[0043] Step 5: Set the total rounds of server-side training and the training target. In each round of server-side training, generate a set number of trainable global prototypes for each category through the trainable prototype module. Minimize the sum of the adaptive margin enhancement contrast loss of the trainable global prototype set of all categories and the clustering prototype set, and minimize the sum of the diversity loss of the trainable global prototype set of all categories. Optimize the parameters of the trainable prototype module to generate trainable global prototype sets of each category that meet the training target;
[0044] Step 5.1: Set the total number of server-side training rounds and training targets, where the training targets include:
[0045] (1) closely align with the cluster prototype set of the mth class to preserve semantic information and maintain a significant distance from the cluster prototype sets of other classes to enhance separability;
[0046] (2) Ensure that the prototypes in the m-th class of trainable global prototypes are orthogonal to each other and maintain the diversity of prototypes within the class;
[0047] Step 5.2: Perform comparative learning on each trainable global prototype and the cluster prototype set obtained in step 4, and introduce the margin value obtained in step 4 into the contrastive loss to achieve close alignment of the trainable global prototype set with the cluster prototype set of the same category and away from the cluster prototype set of different categories;
[0048] Compute the sum of the adaptive margin enhancement contrast loss for all categories and the trainable global prototypes and clustered prototype sets in each category Among them, the adaptive margin enhancement contrast loss between the i-th prototype in the m-th class trainable global prototype set and the clustered prototype set of all categories is As shown in the following formula:
[0049]
[0050] Among them, m′∈[M], m′≠m, m′ is the clustering prototype other than class m, is the Euclidean distance between the ith prototype in the mth class trainable global prototype set and the prototype c in the mth class clustering prototype set, is the Euclidean distance between the ith prototype in the mth class trainable global prototype set and the prototype d in the m′ class clustering prototype set, is the i-th prototype in the m-th class of trainable global prototypes, δ(t) is the margin value when the client and server interact with each other in the t-th round;
[0051] Step 5.3: Calculate the sum of the squares of the cosine similarities between all prototypes in the same type of trainable global prototype set to achieve orthogonalization of the prototypes within the class, so that the trainable global prototypes of the same category can cover richer feature information;
[0052] In order to further increase the dispersion of trainable global prototypes of the same category and prevent the prototypes from being too concentrated in the feature space, for each category, the sum of the squares of the cosine similarities between all global trainable prototypes under the category is calculated as the diversity loss of the trainable prototype set of each category. The sum of the diversity losses of all categories of trainable global prototype sets is Among them, the mth class can train the global prototype set diversity loss As shown in the following formula:
[0053]
[0054] in, is the ith trainable global prototype in the mth class of trainable global prototypes, is the jth trainable global prototype in the mth class of trainable global prototypes, |g m | is the number of trainable global prototypes of the mth class, ||·|| is the modular operation;
[0055] Step 5.4: The total loss is obtained by weighted combination of prototype diversity loss and adaptive margin enhanced contrastive learning loss As shown in the following formula:
[0056]
[0057] With the goal of minimizing the total loss, the parameters of the trainable prototype module are updated based on the gradient descent method, so that the generated prototype can accurately describe the characteristics of the category in the feature space;
[0058] Step 6: Loop step 5 until the total round of server-side training is reached, obtain the optimized trainable prototype module, generate the optimized trainable global prototype set, and distribute the optimized global model parameters and the optimized trainable global prototype set to each client;
[0059] Step 7: Loop steps 2 to 6 until the preset total number of client-server interactions is reached or the global model has converged, and obtain the trained global model. Download the trained global model to the target domain, use the data of the target domain to test the global model, and evaluate the accuracy of the global model in the target domain.
[0060] Step 8: Apply the global model trained by the federated domain generalization method based on trainable prototypes to the image recognition field to complete the corresponding classification task.
[0061] The beneficial effects of adopting the above technical solution are: the federated domain generalization method based on trainable prototypes provided by the present invention optimizes the trainable prototype module through adaptive marginal enhancement contrast learning and prototype diversity learning, improves the inter-class separation and intra-class consistency of the trainable prototype set, so as to improve the generalization performance and classification accuracy of the model in unknown fields. The specific key points are as follows:
[0062] Compared with the prior art, the present invention has the following beneficial effects:
[0063] 1. The present invention designs a trainable prototype module, combines adaptive margin-enhanced contrast learning with prototype diversity learning methods, and optimizes the parameters of the module. During the training process, the trainable prototype module can dynamically adjust the representation of the prototype, thereby balancing the distance within and between classes, and finally generating a prototype set with clear inter-class boundaries and rich intra-class diversity. This strategy significantly improves the classification accuracy, especially when facing heterogeneous data and domain drift, and can effectively improve the performance of the model;
[0064] 2. The present invention introduces dynamic margin values in the comparative learning process, calculates and compares the Euclidean distances between trainable prototypes and prototypes of the same and different categories in the clustering prototype set, prompting the trainable prototype module to reduce the distance between the trainable prototype and the same clustering prototype set during the training process, while widening the distance with the clustering prototype set of different categories, effectively enhancing inter-class separability and maintaining semantic consistency. In addition, for each category, the present invention quantifies the prototype diversity loss by calculating the square sum of cosine similarities between similar trainable global prototypes, improves the diversity of intra-class features, and enables the trainable global prototypes of the same category to cover richer feature information, thereby enhancing the model's generalization ability for different data distributions and complex changes.
[0065] 3. The present invention makes up for the deficiency of existing federated learning methods that do not fully consider the domain generalization problem. It not only ensures data privacy, but also improves the adaptability and accuracy of the model on data in different domains, and significantly improves the generalization performance of the federated learning model in unknown fields. BRIEF DESCRIPTION OF THE DRAWINGS
[0066] Figure 1 A flow chart of a federated domain generalization method based on a trainable prototype provided by an embodiment of the present invention;
[0067] Figure 2 A schematic diagram of data interaction between a client and a server provided in an embodiment of the present invention;
[0068] Figure 3 A flowchart of local training of a client provided in an embodiment of the present invention. DETAILED DESCRIPTION
[0069] The specific implementation of the present invention is further described in detail below in conjunction with the accompanying drawings and examples. The following examples are used to illustrate the present invention, but are not intended to limit the scope of the present invention.
[0070] A federated domain generalization method based on a trainable prototype in this embodiment involves data interaction between a server and several clients, such as Figure 1 As shown, the following steps are included:
[0071] Step 1: Establish a trainable prototype module to generate and update the prototype representation of the category. The initialization parameters of the trainable prototype module include: the number of categories in the dataset, the number of prototypes for each category, the feature dimension, and the hidden layer dimension;
[0072] The trainable prototype module structure includes an embedding layer and a multi-layer nonlinear fully connected network;
[0073] The embedding layer is used to generate an initial embedding representation of the category feature prototype. According to the number of categories and the number of prototypes of each category in the model parameters, the combination of each category and the prototype within the class is mapped to a unique index. The index is used as the input of the embedding layer to generate a corresponding embedding vector. The dimension of each generated embedding vector is determined by the feature dimension specified in the model parameters. The embedding layer uses a normal distribution to initialize the embedding weights, defines the initial distribution of the embedding vector, and makes the embedding vector have good initial separation in the feature space.
[0074] The multi-layer nonlinear fully connected network includes two hidden layers and one output layer; the output feature dimension of the hidden layer is determined by the hidden layer dimension specified in the model parameters, and the hidden layer performs nonlinear mapping on the embedded vector through two fully connected transformations and a ReLU activation function to gradually enhance the expressive power of the feature, and finally, the adjusted feature is mapped to the final prototype representation through the output layer;
[0075] Step 2: If the server interacts with each client for the first time, the server sends the initialized global model parameters to each client. If it is not the first time for the server to interact with each client, the server sends the global model parameters and the global prototype set to each client.
[0076] There is interactive data between the client and the server, such as Figure 2 As shown in the figure, the data distribution between the source domain and the target domain in the client is different. The data in the source domain in the client all carry labels, while the data in the target domain is unlabeled, and the data in the target domain does not participate in the training process;
[0077] Step 3: If the client and the server are interacting for the first time, the client receives the initialized global model parameters sent by the server and initializes the local model. If the client and the server are not interacting for the first time, the client receives the global model parameters and the global prototype set sent by the server. The client uses its own source domain data and the received global prototype set to iteratively train the local model, and uploads the trained local model parameters and local prototype set to the server.
[0078] All clients share a model with the same structure, which includes two modules: feature extractor h and classifier f. The local training process of the client is as follows: Figure 3 As shown, the following steps are included:
[0079] Step 3.1: During the local training process of client k, use the feature extractor h of the local model of client k k For the sample x in the client's own source domain data i Perform feature extraction and obtain the corresponding feature vector z i , the feature vector z i Input to the classifier f and calculate the classifier prediction probability distribution δ(f(z i )) and the corresponding true category y i The cross entropy loss function value of ;
[0080] Feature Extractor h k : The source domain data input by client k Extract feature maps to output space Among them, R V represents the input space of dimension V, is the source domain data of client k, is the source domain data of client k Mapped output space; get the corresponding eigenvector Among them, R D represents the output space with dimension D;
[0081] Classifier f: The feature vector z of client k i Mapped to M-dimensional vector space, generating an unnormalized prediction score f(z i ), by calculating the unnormalized prediction score f(z i ) Apply the softmax function to obtain the predicted probability distribution of each category δ(f(z i )) and calculate the classifier prediction probability distribution δ(f(z i )) and the corresponding true category label y i The cross entropy loss function value is as shown in the following formula:
[0082]
[0083] in, is the true category label y i The corresponding one-hot encoding vector has a dimension that matches the number of categories M. If y i =m, then The mth category of is taken as 1, and the rest are 0;
[0084] Step 3.2: Determine the global prototype set Is it empty? If the global prototype set If empty, set the global prototype contrast learning loss Is 0; if the global prototype set If not empty, the feature extractor h of the local model of client k is calculated. k The feature vector z of the extracted source domain sample i Global prototype comparison learning loss with the global prototype set By minimizing the global prototype contrastive learning loss Close the feature vector z of the source domain sample of client k i The similarity between the global prototype set of the same category and the feature vector z of the source domain sample of client k is pulled away i Similarity with the global prototype set of other categories, using the minimized global prototype contrastive learning loss Optimize the local feature extractor h of client k k , making the feature vector of the client k source domain samples closer to the local prototype of the correct category;
[0085] Global prototype contrastive learning loss As shown below:
[0086]
[0087] Among them, s(z i ,a) is the feature vector z of client k i With the category prototype a∈g m The similarity measure of s(z i ,b) is the feature vector z of client k i With the category prototype b∈g m′ Similarity measure, g is the set of all category prototypes trained on the server side, g m is the prototype set belonging to category m, g m′ =gg m is the set of prototypes that do not belong to category m among all category prototypes;
[0088] Step 3.3: Compare the global prototype to the learning loss l GPCL and cross entropy loss The weighted combination is the total loss, minimizing the total loss and optimizing the client local model parameter θ k ;
[0089] Step 3.4: Loop through steps 3.1 to 3.3 until the total number of local training rounds on the client k is reached. k The output feature vector z i As the class prototype information, the average feature vector of all samples in category m is represented as the prototype of the mth class, and the local prototype set O is updated. k ;
[0090] Take the average of the feature vectors of the same category on client k to obtain the local prototype of category m on client k As shown in the following formula:
[0091]
[0092] in, is the dataset whose label belongs to category m in client k, is the j-th sample belonging to category m in client k, For sample x j The corresponding label, h k (x j ) is the local feature extractor h on client k k For sample x j The extracted feature vector, k = {1, 2, ..., K} is the client set, m = {1, 2, ..., M} is the category set, M is the number of categories, is the local prototype of category m on client k;
[0093] Then we get the local prototype set O of all categories of client k k , as shown in the following formula:
[0094]
[0095] Step 3.5: Upload client k’s local prototype set O k and local model parameters θ k Go to the server and wait for the next data interaction with the server;
[0096] Step 4: The server receives the local prototype sets and local model parameters uploaded by each client, performs weighted average of the local model parameters of each client according to the proportion of the local data volume of each client to the total data volume, and generates updated global model parameters; uses an unsupervised clustering algorithm to cluster the local prototype sets belonging to the same category to obtain clustered prototype sets of the same category; calculates the minimum Euclidean distance between the clustered prototype sets of each category, obtains the maximum inter-class distance between all categories, compares it with the pre-set threshold, and selects the smaller value of the two as the margin value for data interaction between the client and the server in this round;
[0097] The updated global model parameter θ′ is:
[0098]
[0099] Among them, N k is the number of samples on client k, and N is the total number of samples on all clients;
[0100] The unsupervised clustering algorithm is based on the nearest neighbor of each local prototype. By calculating the cosine similarity between prototypes, similar prototypes are automatically classified into the same clustering result. The unsupervised clustering algorithm is used to cluster the local prototype sets belonging to the same category. According to the clustering results, several representative prototypes are selected as clustering prototypes for each category. The specific method for obtaining the clustering prototype set C of all categories is as follows:
[0101]
[0102] C={C 1 ,...,C m ,...C M}
[0103] Among them, C is the clustering prototype set of all categories, C m is the clustering prototype set of the mth class, For N m The representative prototype obtained by clustering the m-th local prototypes, J m is the number of cluster prototypes of the mth class, N m is the number of local prototypes of the mth type, Cluster is an unsupervised clustering algorithm operation, and the input m-type local prototype set O k Clustering is performed together to obtain J m Cluster centers;
[0104] By comparing with the pre-set threshold, the margin value is prevented from exceeding the reasonable range, making the training process more stable;
[0105] Step 5: Set the total rounds of server-side training and the training target. In each round of server-side training, generate a set number of trainable global prototypes for each category through the trainable prototype module. Minimize the sum of the adaptive margin enhancement contrast loss of the trainable global prototype set of all categories and the clustering prototype set, and minimize the sum of the diversity loss of the trainable global prototype set of all categories. Optimize the parameters of the trainable prototype module to generate trainable global prototype sets of each category that meet the training target;
[0106] Step 5.1: Set the total number of server-side training rounds and training targets, where the training targets include:
[0107] (1) closely align with the cluster prototype set of the mth class to preserve semantic information and maintain a significant distance from the cluster prototype sets of other classes to enhance separability;
[0108] (2) Ensure that the prototypes in the m-th class of trainable global prototypes are orthogonal to each other and maintain the diversity of prototypes within the class;
[0109] Step 5.2: Perform comparative learning on each trainable global prototype and the cluster prototype set obtained in step 4, and introduce the margin value obtained in step 4 into the contrastive loss to achieve close alignment of the trainable global prototype set with the cluster prototype set of the same category and away from the cluster prototype set of different categories;
[0110] Compute the sum of the adaptive margin enhancement contrast losses for all categories and the trainable global prototypes and clustered prototype sets in each category Among them, the adaptive margin enhancement contrast loss between the i-th prototype in the m-th class trainable global prototype set and the clustered prototype set of all categories is As shown in the following formula:
[0111]
[0112] Among them, m′∈[M], m′≠m, m′ is the clustering prototype other than class m, is the Euclidean distance between the ith prototype in the mth class trainable global prototype set and the prototype c in the mth class clustered prototype set, is the Euclidean distance between the ith prototype in the mth class trainable global prototype set and the prototype d in the m′ class clustering prototype set, is the i-th prototype in the m-th class of trainable global prototypes, δ(t) is the margin value when the client and server interact with each other in the t-th round;
[0113] Step 5.3: Calculate the sum of the squares of the cosine similarities between all prototypes in the same type of trainable global prototype set to achieve orthogonalization of the prototypes within the class, so that the trainable global prototypes of the same category can cover richer feature information;
[0114] In order to further increase the dispersion of trainable global prototypes of the same category and prevent the prototypes from being too concentrated in the feature space, for each category, the sum of the squares of the cosine similarities between all global trainable prototypes under the category is calculated as the diversity loss of the trainable prototype set of each category. The sum of the diversity losses of all categories of trainable global prototype sets is Among them, the mth class can train the global prototype set diversity loss As shown in the following formula:
[0115]
[0116] in, is the ith trainable global prototype in the mth class of trainable global prototypes, is the jth trainable global prototype in the mth class of trainable global prototypes, |g m | is the number of trainable global prototypes of the mth class, ||·|| is the modular operation;
[0117] Step 5.4: The total loss is obtained by weighted combination of prototype diversity loss and adaptive margin enhanced contrastive learning loss As shown in the following formula:
[0118]
[0119] With the goal of minimizing the total loss, the parameters of the trainable prototype module are updated based on the gradient descent method, so that the generated prototype can accurately describe the characteristics of the category in the feature space;
[0120] Step 6: Loop step 5 until the total round of server-side training is reached, obtain the optimized trainable prototype module, generate the optimized trainable global prototype set, and distribute the optimized global model parameters and the optimized trainable global prototype set to each client;
[0121] Step 7: Loop steps 2 to 6 until the preset total number of client-server interactions is reached or the global model has converged, and obtain the trained global model. Download the trained global model to the target domain, use the data of the target domain to test the global model, and evaluate the accuracy of the global model in the target domain.
[0122] Step 8: Apply the global model trained by the federated domain generalization method based on trainable prototypes to the image recognition field to complete the corresponding classification task.
[0123] Compared with the prototype-based federated learning method, the global model trained by the federated domain generalization method based on trainable prototypes provided in this embodiment exhibits stronger generalization ability when processing samples with unknown data distribution, especially when facing large differences in data distribution.
[0124] During the training process, in view of the potential distribution differences of data in different fields, the federated domain generalization method based on trainable prototypes provided in this embodiment introduces adaptive marginal enhancement contrast learning and prototype diversity learning technology, optimizes the trainable prototype module, improves the inter-class separation and intra-class consistency of the prototype set, and enhances the adaptability of the global model to different data distributions and complex changes. In the face of unknown data distribution samples, the global model trained by the federated domain generalization method based on trainable prototypes can provide more accurate classification, thereby effectively alleviating the overfitting problem in traditional methods, significantly improving the prediction accuracy and stability in different fields and scenarios, and ultimately providing excellent performance for practical applications.
[0125] Comparative experiments were conducted on the PACS, OfficeHome, and OfficeCaltech datasets using the federated domain generalization method based on a trainable prototype provided in this embodiment and other federated domain generalization methods. The results of the comparative experiments are shown in Table 1.
[0126] Table 1 Comparative experimental results
[0127]
[0128] As can be seen from Table 1, the federated domain generalization method FedTCP based on trainable prototypes provided in this embodiment is superior to traditional methods in accuracy, generalization performance and consistency between different clients; on PACS, compared with the traditional FedAvg method, the accuracy of this solution is improved by 5.5%, compared with the best performing federated domain generalization method FedSR, the accuracy is improved by 2.15%, and compared with the prototype-based federated learning method FPL, the accuracy is improved by 0.6%; on OfficeHome, the accuracy of this solution is improved by 0.62% compared with FPL, and on OfficeCaltech, the accuracy of this solution is improved by 0.95% compared with FPL, achieving higher generalization and discrimination capabilities with relatively low communication costs, providing an effective solution for solving federated learning under domain generalization;
[0129] Finally, it should be noted that the above embodiments are only used to illustrate the technical solutions of the present invention, rather than to limit it. Although the present invention has been described in detail with reference to the aforementioned embodiments, those skilled in the art should understand that they can still modify the technical solutions described in the aforementioned embodiments, or make equivalent replacements for some or all of the technical features therein. However, these modifications or replacements do not cause the essence of the corresponding technical solutions to deviate from the scope defined by the claims of the present invention.
Claims
1. A federated domain generalization method based on trainable prototypes, characterized by: Involving data interaction between a server and several clients, including the following steps: Step 1: Establish a trainable prototype module to generate and update the prototype representation of the category. The initialization parameters of the trainable prototype module include: the number of categories in the dataset, the number of prototypes for each category, the feature dimension, and the hidden layer dimension; Step 2: If the server and the client are interacting for the first time, the server sends the initialized global model parameters to the client. If the server and the client are not interacting for the first time, the server sends the global model parameters and the global prototype set to the client. Step 3: If the client and the server are interacting for the first time, the client receives the initialized global model parameters sent by the server and initializes the local model. If the client and the server are not interacting for the first time, the client receives the global model parameters and the global prototype set sent by the server. The client uses its own source domain data and the received global prototype set to iteratively train the local model, and uploads the trained local model parameters and local prototype set to the server. Step 4: The server receives the local model parameters and local prototype sets uploaded by each client, performs weighted average of the local model parameters of each client according to the proportion of the local data volume of each client to the total data volume, and generates updated global model parameters; uses an unsupervised clustering algorithm to cluster the local prototype sets belonging to the same category to obtain clustered prototype sets of the same category; calculates the minimum Euclidean distance between the clustered prototype sets of each category, obtains the maximum inter-class distance between all categories, compares it with the pre-set threshold, and selects the smaller value of the two as the margin value for data interaction between the client and the server in this round; Step 5: Set the total rounds of server-side training and the training target. In each round of server-side training, generate a set number of trainable global prototypes for each category through the trainable prototype module. Minimize the sum of the adaptive margin enhancement contrast loss of the trainable global prototype set of all categories and the clustering prototype set, and minimize the sum of the diversity loss of the trainable global prototype set of all categories. Optimize the parameters of the trainable prototype module to generate trainable global prototype sets of each category that meet the training target; Step 6: Loop step 5 until the total round of server-side training is reached, obtain the optimized trainable prototype module, generate the optimized trainable global prototype set, and distribute the optimized global model parameters and the optimized trainable global prototype set to each client; Step 7: Loop steps 2 to 6 until the preset total number of client-server interactions is reached or the global model has converged, and obtain the trained global model. Download the trained global model to the target domain, use the data of the target domain to test the global model, and evaluate the accuracy of the global model in the target domain. Step 8: Apply the global model trained by the federated domain generalization method based on trainable prototypes to the image recognition field to complete the corresponding classification task.
2. The method for generalizing federated domains based on trainable prototypes according to claim 1, characterized in that: In the step 1, the trainable prototype module structure includes an embedding layer and a multi-layer nonlinear fully connected network; The embedding layer is used to generate an initial embedding representation of the category feature prototype; according to the number of categories and the number of prototypes of each category in the model parameters, the combination of each category and the prototype within the class is mapped to a unique index, and the index is used as the input of the embedding layer to generate a corresponding embedding vector, and the dimension of each generated embedding vector is determined by the feature dimension specified in the model parameters; The embedding layer uses normal distribution to initialize the embedding weights and define the initial distribution of the embedding vectors, so that the embedding vectors have good initial separation in the feature space; The multi-layer nonlinear fully connected network includes two hidden layers and one output layer; the output feature dimension of the hidden layer is determined by the hidden layer dimension specified in the model parameters, and the hidden layer performs nonlinear mapping on the embedded vector through two fully connected transformations and a ReLU activation function to gradually enhance the expressive power of the features, and finally, the adjusted features are mapped to the final prototype representation through the output layer.
3. The method for generalizing federated domains based on trainable prototypes according to claim 2, characterized in that: In step 3, all clients share a model with the same structure, which includes two modules: a feature extractor h and a classifier f.
4. The method for generalizing federated domains based on trainable prototypes according to claim 3, characterized in that: The local training process of the client in step 3 includes the following steps: Step 3.1: During the local training process of client k, use the feature extractor h of the local model of client k k For the sample x in the client's own source domain data i Perform feature extraction and obtain the corresponding feature vector z i , the feature vector z i Input to the classifier f and calculate the classifier prediction probability distribution δ(f(z i )) and the corresponding true category y i The cross entropy loss function value of ; Feature Extractor The source domain data input by client k Extract feature maps to output space Among them, R V represents the input space of dimension V, is the source domain data of client k, is the source domain data of client k Mapped output space; get the corresponding eigenvector Among them, R D represents the output space with dimension D; Classifier The feature vector z of client k i Mapped to M-dimensional vector space, generating an unnormalized prediction score f(z i ), by calculating the unnormalized prediction score f(z i ) Apply the softmax function to obtain the predicted probability distribution of each category δ(f(z i )) and calculate the classifier prediction probability distribution δ(f(z i )) and the corresponding true category label y i The cross entropy loss function value is as shown in the following formula: in, is the true category label y i The corresponding one-hot encoding vector has a dimension that matches the number of categories M. If y i =m, then The mth category of is taken as 1, and the rest are 0; Step 3.2: Determine the global prototype set Is it empty? If the global prototype set If it is empty, the global prototype contrast learning loss e is set. GPCL Is 0; if the global prototype set If not empty, the feature extractor h of the local model of client k is calculated. k The feature vector z of the extracted source domain sample i The global prototype comparison learning loss e with the global prototype set GPCL , by minimizing the global prototype contrastive learning loss e GPCL Close the feature vector z of the source domain sample of client k i The similarity between the global prototype set of the same category and the feature vector z of the source domain sample of client k is pulled away i Similarity with the global prototype set of other categories, using the minimized global prototype contrastive learning loss e GPCL Optimize the local feature extractor h of client k k , making the feature vector of the client k source domain samples closer to the local prototype of the correct category; Global prototype contrastive learning loss e GPCL As shown below: Among them, s(z i ,a) is the feature vector z of client k i With the category prototype a∈g m The similarity measure of s(z i ,b) is the feature vector z of client k i With the category prototype b∈g m′ Similarity measure, g is the set of all category prototypes trained on the server side, g m is the prototype set belonging to category m, g m′ =gg m is the set of prototypes that do not belong to category m among all category prototypes; Step 3.3: Compare the global prototype to the learning loss l GPCL and the cross entropy loss l CE The weighted combination is the total loss, minimizing the total loss and optimizing the client local model parameter θ k ; Step 3.4: Loop through steps 3.1 to 3.3 until the total number of local training rounds on the client k is reached. k The output feature vector z i As the class prototype information, the average feature vector of all samples in category m is represented as the prototype of the mth class, and the local prototype set O is updated. k ; Take the average of the feature vectors of the same category on client k to obtain the local prototype of category m on client k As shown in the following formula: in, is the dataset whose label belongs to category m in client k, is the j-th sample belonging to category m in client k, For sample x j The corresponding label, h k (x j ) is the local feature extractor h on client k k For sample x j The extracted feature vector, k = {1, 2, ..., K} is the client set, m = {1, 2, ..., M} is the category set, M is the number of categories, is the local prototype of category m on client k; Then we get the local prototype set O of all categories of client k k , as shown in the following formula: Step 3.5: Upload client k’s local prototype set O k and local model parameters θ k To the server, waiting for the next data interaction with the server.
5. The method for generalizing federated domains based on trainable prototypes according to claim 4, characterized in that: The global model parameter θ′ after the update in step 4 is: Among them, N k is the number of samples on client k, and N is the total number of samples on all clients; According to the clustering results, several representative prototypes are selected as clustering prototypes for each category. The specific method to obtain the clustering prototype set C of all categories is as follows: C={C 1 ,...,C m ,...C M } Among them, C is the clustering prototype set of all categories, C m is the clustering prototype set of the mth class, For N m The representative prototype obtained by clustering the m-th local prototypes, J m is the number of cluster prototypes of the mth class, N m is the number of local prototypes of the mth type, Cluster is an unsupervised clustering algorithm operation, and the input m-type local prototype set O k Clustering is performed together to obtain J m Cluster centers.
6. The method for generalizing federated domains based on trainable prototypes according to claim 5, characterized in that: The step 5 specifically includes: Step 5.1: Set the total number of server-side training rounds and training targets; Step 5.2: Perform comparative learning on each trainable global prototype and the cluster prototype set obtained in step 4, and introduce the margin value obtained in step 4 into the contrastive loss to achieve close alignment of the trainable global prototype set with the cluster prototype set of the same category and away from the cluster prototype set of different categories; Compute the sum of the adaptive margin enhancement contrast losses for all categories and the trainable global prototypes and clustered prototype sets in each category Among them, the adaptive margin enhancement contrast loss between the i-th prototype in the m-th class trainable global prototype set and the clustered prototype set of all categories is As shown in the following formula: Among them, m′∈[M], m′≠m, m′ is the clustering prototype other than class m, is the Euclidean distance between the ith prototype in the mth class trainable global prototype set and the prototype c in the mth class clustered prototype set, is the Euclidean distance between the ith prototype in the mth class trainable global prototype set and the prototype d in the m′ class clustering prototype set, is the i-th prototype in the m-th class of trainable global prototypes, δ(t) is the margin value when the client and server interact with each other in the t-th round; Step 5.3: Calculate the sum of the squares of the cosine similarities between all prototypes in the same type of trainable global prototype set to achieve orthogonalization of the prototypes within the class, so that the trainable global prototypes of the same category can cover richer feature information; In order to further increase the dispersion of trainable global prototypes of the same category and prevent the prototypes from being too concentrated in the feature space, for each category, the sum of the squares of the cosine similarities between all global trainable prototypes under the category is calculated as the diversity loss of the trainable prototype set of each category. The sum of the diversity losses of all categories of trainable global prototype sets is Among them, the mth class can train the global prototype set diversity loss l m As shown in the following formula: in, is the ith trainable global prototype in the mth class of trainable global prototypes, is the jth trainable global prototype in the mth class of trainable global prototypes, |g m | is the number of trainable global prototypes of the mth class, ||·|| is the modular operation; Step 5.4: According to the weighted combination of prototype diversity loss and adaptive margin enhanced contrastive learning loss, the total loss l is obtained as shown in the following formula: With the goal of minimizing the total loss, the parameters of the trainable prototype module are updated based on the gradient descent method, so that the generated prototype can accurately describe the characteristics of the category in the feature space.
7. The method for generalizing federated domains based on trainable prototypes according to claim 6, characterized in that: The training objectives of step 5.1 include: (1) closely align with the cluster prototype set of the mth class to preserve semantic information and maintain a significant distance from the cluster prototype sets of other classes to enhance separability; (2) Ensure that the prototypes in the m-th class of trainable global prototypes are orthogonal to each other and maintain the diversity of prototypes within the class.
Citation Information
Cited By
Personalized federal learning method and system based on domain invariant text representation and intra-domain global prior
CN120508883A
Federal learning method and system based on domain-invariant text representation and global prior in domain
CN120508883B
Global training method and local training method based on federal learning
CN121413711A