Personalized federated learning method for asynchronous optimization and prototype perception reasoning of classifier
Through the method of asynchronous optimization of classifiers and prototype-aware reasoning, the problems of label offset and domain offset in personalized federated learning are solved, and the coordinated optimization of personalization and generalization performance is achieved under the premise of protecting privacy. It is suitable for personalized model optimization tasks in medical image analysis.
Patent Information
- Application Number
- CN202510871190.2
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-06-26
- Publication Date
- 2025-09-23
- Estimated Expiration
- 2045-06-26
AI Technical Summary
Existing personalized federated learning methods have privacy leakage risks when dealing with label shift and domain shift, and have insufficient performance in weakly heterogeneous environments, making it difficult to achieve effective knowledge transfer and personalization needs.
The method of asynchronous classifier optimization and prototype-aware reasoning is adopted. By designing an asynchronously updated dual classifier mechanism in the training phase and introducing a bilateral prototype clustering strategy of local prototypes and global prototypes, the output weights of the two classifiers are adaptively calculated. In the reasoning phase, prototype-aware technology is used to achieve the coordinated optimization of personalization ability and generalization performance.
Under the premise of protecting data privacy, it effectively alleviates the impact of label offset and domain offset, improves the personalization and generalization performance of the model in different heterogeneous environments, reduces communication costs, and is suitable for personalized model optimization tasks in medical image analysis.
Smart Images

Figure CN120687850A_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the field of federated learning technology, and in particular relates to a personalized federated learning method for asynchronous optimization of classifiers and prototype-aware reasoning. Background Art
[0002] Federated learning, a distributed machine learning paradigm, was originally designed to address the "data silo" problem and facilitate multi-party data collaboration without the need for direct exchange of raw data. With the continuous evolution of technology and the rapid growth of the user base, a single, generalized global model is no longer able to meet the personalized service needs of each client. To this end, personalized federated learning has emerged. While sharing global knowledge, personalized federated learning optimizes local model performance based on the specific data distribution and task requirements of each client. Compared to traditional federated learning, personalized federated learning not only effectively alleviates the challenges of data heterogeneity but also significantly improves the performance of local models on specific clients with minimal sacrifice in global generalization capabilities.
[0003] As an efficient and practical technical approach, model decoupling is widely used in personalized federated learning. A typical approach is to divide the model into two parts: a feature extractor and a classifier. The feature extractor is uploaded to the server for aggregation, while the classifier is retained locally to learn personalized features. This approach not only improves the model's personalization capabilities but also preserves rich shared knowledge. Therefore, it has become a research hotspot in the field of personalized federated learning in recent years. However, due to the heterogeneity of client data distribution, simple feature extractor average aggregation strategies often fail to achieve effective knowledge transfer and may even lead to negative transfer, causing performance degradation. Therefore, maintaining the semantic consistency of features across clients becomes the key to achieving effective knowledge transfer. Currently, the mainstream approach is to guide feature extractor updates based on feature prototypes to encourage each client to learn globally consistent feature representations.
[0004] Although existing personalized federated learning methods based on model decoupling perform well in strongly heterogeneous environments, they generally have the problem of insufficient adaptability to weakly heterogeneous environments, especially when dealing with label shifts. For example, FedPer and FedRep can effectively improve performance in strongly heterogeneous environments by uploading feature extractors and retaining local classifiers, but under weakly heterogeneous conditions, the model performance has declined significantly. In the direction of feature alignment, methods such as FedProto propose uploading only feature prototypes instead of model parameters, and constructing a global prototype through an averaging strategy. However, when the category distribution of the client is completely different or there are serious category omissions, the expressive power of the global prototype is limited and it cannot accurately cover all local feature distributions, resulting in a decline in generalization performance and difficulty in effectively alleviating the impact of domain shift.
[0005] Chinese patent document CN118734995A discloses a single-client multi-domain heterogeneous federated learning system and method based on manifold learning. Each client determines multiple manifold point sets corresponding to the categories of data in the training set based on a local model and a local dataset, and sends the manifold point sets and local models to a server. The server determines a global manifold based on the manifold points of each client in any category. It partitions the data domain of each global manifold into submanifolds through clustering. The server updates each submanifold using an attention mechanism, reconstructs the manifold points, and updates the global model based on each local model. The reconstructed manifold points and the current global model are distributed to the corresponding client. Each client determines a total loss function based on the received global model, updates the current local model based on the total loss function, and sends the updated local model to the server. The client and server iteratively train until the global model converges, thereby achieving good federated performance in domain heterogeneous scenarios. However, this method uploads the entire model and manifold point set, which incurs high communication costs and poses a risk of privacy leakage. This method only addresses the problem of domain generalization but ignores the need for client personalization, failing to effectively balance the needs of client generalization and personalization.
[0006] Chinese patent document CN119658639A discloses a Non-IID federated learning method and device based on trajectory matching dataset distillation, which includes: Step 1: There are N participants, each participant has a local data distribution and label, and the data distribution between clients does not satisfy independent and identical distribution; on the client, each client initializes its local model and replaces the original classifier of the given model architecture with the same ETF classifier; the server runs a standard federated averaging algorithm, the client trains its model on local data, updates the model parameters, and uploads the parameters to the server, the server aggregates the parameters and updates the global model, and this process is repeated continuously; the server initializes a latent vector set of synthetic data and optimizes it according to the early trajectory of the global model; the server uses the optimized synthetic data to correct the aggregated global model until the number of communication rounds reaches rounds, initializes an extended latent vector set, and continues to optimize based on the later global model trajectory; the global model is further corrected using the synthetic data generated by the first two optimizations; with the continuous updating, aggregation and correction of the model, the global model gradually converges. However, it is only based on latent vector synthesis and does not fully utilize global semantic consistency, which will lead to the loss of semantic features. The model is difficult to learn shared knowledge across distributions during aggregation, resulting in a decrease in the predictive ability of unknown distribution data. In addition, it may require more local data to participate in training to make up for the semantic loss, thereby increasing the data transmission volume and the risk of privacy leakage. It mainly focuses on the alignment of feature space and does not adequately handle label offsets.
[0007] In recent years, several improvements have been proposed to address these issues. For example, FedPAC optimizes by combining feature prototype alignment with classifier collaboration. First, feature prototypes are used to guide feature extractor updates, mitigating domain shift; second, the server-side classifier is restructured based on client distribution similarities. This strategy demonstrates good personalization and generalization performance across a variety of heterogeneous environments. However, FedPAC requires clients to upload local data distribution information, which introduces potential privacy risks.
[0008] Therefore, there is an urgent need to explore a new personalized federated learning method that does not require exposing private information, takes into account both label shift and domain shift, and can achieve efficient personalization and generalization performance. Summary of the Invention
[0009] The present invention aims to overcome the defects of the above-mentioned prior art and proposes a personalized federated learning method of asynchronous optimization of classifiers and prototype-aware reasoning. By designing an asynchronously updated dual-classifier mechanism in the training phase and utilizing a prototype-aware technology in the reasoning phase, the output weights of the two classifiers are adaptively calculated, so that the model can adaptively judge the confidence of personalized knowledge and generalized knowledge, thereby achieving dynamic complementarity and balance between the two. In addition, in order to ensure that the features generated by each client feature extractor have a globally consistent semantic expression, the present invention introduces a bilateral prototype clustering strategy of local prototypes and global prototypes to generate a unified category prototype and adaptively guide each client feature extractor to update in a consistent direction, thereby effectively alleviating the performance degradation problem caused by domain shift. By combining asynchronous update of classifiers with bilateral prototype clustering technology, the present invention achieves the coordinated optimization of personalization capabilities and generalization performance under the premise of protecting data privacy, and effectively responds to the challenges of multi-client collaboration under label shift and domain shift environments.
[0010] In order to solve the above technical problems, the present invention provides a personalized federated learning method for asynchronous optimization of classifiers and prototype-aware reasoning, the method comprising: Training phase: S0, server-side initializes shared classifier and global clustering prototypes , initialize the shared classifier and global clustering prototypes Broadcast to each client participating in the training; S1. Asynchronous update: The client receives the shared classifier and global clustering prototype from the server. The client first freezes the personalized classifier, uses the global clustering prototype to adaptively align features, and guides the update of the personalized feature extractor and shared classifier. After the update, the personalized classifier is unfrozen. The client then freezes the updated personalized feature extractor and shared classifier, updates the personalized classifier, and unfreezes them after the update. S2. Use the updated personalized feature extractor to extract features from the local training set, and perform local prototype clustering on the extracted features to obtain the local cluster prototype set. , and then perform weighted average calculation on the prototypes in the set to obtain the local unbiased prototype ; S3, the client uses the shared classifier updated in step S1 and the local unbiased prototype obtained in step S2 Send to server; S4. The server aggregates the received shared classifiers using the average aggregation method to obtain the aggregated shared classifiers. ; S5. The server receives the local unbiased prototype Perform global prototype clustering to obtain the global cluster prototype ; S6. The server uses the aggregated shared classifier obtained in step S4 and the global clustering prototype obtained in step S5 Rebroadcast to all participating clients and repeat steps S1-S6 until the preset rounds are reached or the model converges; Adaptive reasoning stage based on prototype perception: S7. After training is completed, the corresponding features of the local test set extracted by the personalized feature extractor are input into the shared classifier and the personalized classifier to obtain personalized prediction output and shared prediction output. Based on prototype perception technology, the client adaptively calculates the output weights of the two classifiers and performs weighted fusion on the prediction outputs of the two classifiers to obtain the final prediction result, and then obtains the final test accuracy.
[0011] Preferably, the step S1 uses the global clustering prototype to adaptively perform feature alignment to guide the update of the personalized feature extractor and the shared classifier, specifically: S11. The client calculates the entropy value based on the label distribution of the local training set to obtain the adaptive alignment weight for feature alignment. : Statistics on the label distribution of each client, and use formula (1) to calculate the entropy of label distribution: (1) In formula (1), represents the entropy of the label distribution of client k, Represents the random variable Y taking the tth category in the label space The probability of represents the t-th category in the label space, C represents the total number of categories in the label space, represent The logarithm of , the larger its value, the more uneven the label distribution, which is used to evaluate the degree of heterogeneity of data distribution; Use the obtained entropy value to calculate the adaptive alignment weight of client k : (2) In formula (2), is the adaptive alignment weight of client k, is the entropy of the label distribution of client k, is the scaling factor; S12, based on adaptive alignment weight , update the personalized feature extractor and shared classifier: For updating the personalized feature extractor, we use the global clustering prototype previously obtained from the server in conjunction with contrastive learning to guide it to generate semantically consistent and class-separable features. That is, features of the same class are close to each other, and features of different classes are far away from each other. The loss of feature alignment is calculated as follows: (3) In formula (3), is the loss function of the personalized feature extractor, The feature representation of the i-th sample extracted by the personalized feature extractor of the k-th client before updating, represent The cosine similarity between and the category prototype c, Represents the category The set of global cluster prototypes is the positive sample set, Representatives do not belong to the category The set of global clustering prototypes is the negative sample set. By calculating this loss, the personalized feature extractor of each client can produce semantically consistent feature representations while maintaining the separability between categories. For the update of the shared classifier, the traditional cross entropy loss is used: (4) In formula (4), is the cross entropy loss of the shared classifier, | | represents the local training set of client K The number of samples, is the label value of the i-th sample, Represents the prediction value of the shared classifier for the i-th sample, and the shared classifier is updated by minimizing the loss; The total loss of the personalized feature extractor and shared classifier update is: (5) In formula (5), is the total loss of the update, is the cross entropy loss of the shared classifier, represents the adaptive weight for feature alignment, is the loss function of the personalized feature extractor, which is minimized by , to improve the performance of the personalized feature extractor and the shared classifier, and obtain the updated personalized feature extractor and the shared classifier; Step S1 of updating the personalized classifier is as follows: Use cross entropy loss to update it, and its loss function is: (6) (7) in, is the total loss of the personalized classifier update, represents the cross entropy loss of the personalized classifier, | | represents the local training set of client K The number of samples, is the label value of the i-th sample, Represents the predicted output of the personalized classifier of the i-th sample; by minimizing To improve the classification performance of the personalized classifier and obtain an updated personalized classifier.
[0012] Preferably, in step S2, the extracted features are clustered into local prototypes to obtain a local cluster prototype set. , specifically: Client k performs FINCH clustering on the features extracted from the personalized feature extractor. The clustering process is as follows: (8) In formula (8), represents the set of local cluster prototypes of client k belonging to category m, Represents the jth feature prototype of client k belonging to category m, represents the number of cluster prototypes of client k belonging to category m, Represents the input sample The features obtained by the updated personalized feature extractor, represents the i-th input sample and its label, represents the local training set of the kth client belonging to category m; The weighted average calculation of the prototypes in the set is performed to obtain a local unbiased prototype , specifically: After obtaining the feature prototype of each category, the feature prototypes of each category are averaged to obtain the local unbiased prototype. The calculation process is as follows: (9) In formula (9), represents the local unbiased prototype of client k belonging to category m, Represents the jth feature prototype of client k belonging to category m, Represents the number of cluster prototypes for which client k belongs to category m.
[0013] Preferably, the step S4 is specifically as follows: The server aggregates the local shared classifier using the following formula: (10) In formula (10), represents the shared classifier after aggregation, Represents the number of clients participating in the aggregation, is the local shared classifier of the k-th client.
[0014] Preferably, in step S5, the server receives the local unbiased prototype Perform global prototype clustering to obtain the global cluster prototype , specifically: Use FINCH to upload local unbiased prototypes Clustering is performed, and the clustering process is as follows: (11) In formula (11), represents the set of global cluster prototypes of category m, represents the t-th global cluster prototype of category m, represents the number of global cluster prototypes for category m, represents the local unbiased prototype of category m for client k, A set of local unbiased prototypes representing each client class m.
[0015] Preferably, the step S7 is specifically as follows: S71. After training is completed, first, the local test set is fed into the personalized feature extractor to extract corresponding features, and the features are fed into the personalized classifier and the shared classifier to obtain personalized prediction output and shared prediction output respectively; S72, calculate the features extracted by the personalized feature extractor and the global prototype and locally unbiased prototypes The cosine similarity of The prototype of global clustering The average calculation is obtained; S73, using the softmax function to normalize the two cosine similarities calculated in step S72 to obtain weights of the prediction outputs of the shared classifier and the personalized classifier, and using the weights to fuse the two prediction outputs obtained in step S71 to obtain the final inference prediction output; S74. Compare the final inference prediction output obtained in step S73 with the actual output result to obtain the final test accuracy.
[0016] Further preferably, the features extracted by the calculation personalized feature extractor in step S72 are respectively compared with the global prototype and locally unbiased prototypes The cosine similarity is: (12) (13) (14) In formula (12), (13), (14), For the client Features generated by the personalized feature extractor during the inference phase and corresponding categories Locally unbiased prototype of The cosine similarity between For the client Features generated by the personalized feature extractor during the inference phase and corresponding categories The global prototype The cosine similarity between yes The standardized measure of yes The standardized measure of yes The standardized measure of represents the number of global cluster prototypes for category m, Represents the tth global cluster prototype of category m.
[0017] Further preferably, the step S73 is specifically as follows: The weight between the prediction output of the shared classifier and the prediction output of the personalized classifier is calculated by formula (15): (15) In formula (15), and represent the weight of the prediction output of the shared classifier and the weight of the prediction output of the personalized classifier for client k, respectively, and + =1, For the client Features generated by the personalized feature extractor and corresponding categories Locally unbiased prototype of The cosine similarity between For the client Features generated by the personalized feature extractor and corresponding categories The global prototype The cosine similarity between According to the weights of the prediction outputs of the personalized classifier and the shared classifier, the final inference prediction output is: (16) in, is the inference prediction output, and represent the prediction output weights of the shared classifier and the prediction output weights of the personalized classifier for client k, respectively. is the prediction output of the personalized classifier for client k, is the prediction output of the shared classifier.
[0018] The present invention also provides an application of a personalized federated learning method of asynchronous optimization of classifiers and prototype-aware reasoning in medical image analysis, which is used for personalized model optimization tasks in medical image analysis.
[0019] Compared with the prior art, the present invention has the following beneficial effects: (1) The method of the present invention designs an asynchronously updated dual classifier mechanism and introduces a bilateral prototype clustering strategy of local prototypes and global prototypes in the training phase. In the inference phase, a prototype-aware technology is used to adaptively calculate the output weights of the two classification heads. This achieves the coordinated optimization of personalization ability and generalization performance while protecting privacy, and alleviates the impact of label shift and domain shift.
[0020] (2) Label shift: Due to the existence of label shift, that is, the data label distribution between clients presents a non-IID characteristic, which leads to inconsistent update directions of local models, thereby significantly reducing the performance of the global model. To alleviate this problem, personalized federated learning methods reduce the impact of label shift to a certain extent by training personalized models for each client. However, the existing personalized federated learning methods still have the problem of insufficient performance in weakly heterogeneous scenarios, so that the label shift problem has not been completely solved. To this end, the present invention introduces a shared classifier, which effectively integrates personalized knowledge and generalized knowledge by adaptively fusing the prediction outputs of two classifiers in the inference stage, thereby showing excellent performance in scenarios with different degrees of heterogeneity, further alleviating and solving the label shift problem; the present invention uses a loss-guided dual classifier mechanism, and the model can not only deal with the label shift problem in strong heterogeneous scenarios through personalized heads, but also alleviate the label shift problem in weakly heterogeneous scenarios through shared heads, thereby solving the label shift problem in various heterogeneous environments.
[0021] (3) Domain shift: When client data comes from different domains, existing personalized federated learning methods are difficult to effectively deal with, resulting in a significant decrease in the model's domain generalization ability. By introducing a bilateral prototype clustering strategy, a global category prototype rich in more diverse knowledge and suitable for each client is obtained. This guides the update of the personalized feature extractor, prompting each client to generate semantically consistent feature representations, thereby effectively alleviating the domain shift phenomenon and improving the model's generalization performance.
[0022] (4) Communication cost: Since the method of the present invention only uploads the shared classifier and the local unbiased prototype, the communication efficiency is greatly improved compared to the method of uploading the entire model.
[0023] (5) Protecting data privacy: Since the method of the present invention does not need to upload any information related to local data distribution to the server, the risk of privacy leakage is avoided.
[0024] (6) The present invention is particularly suitable for personalized model optimization tasks in medical imaging analysis, such as scenarios where data distribution varies between different hospitals, different imaging devices, or different patient groups. Traditional federated learning methods, due to their lack of effective adaptation to local data characteristics and label offsets, have difficulty in simultaneously addressing individualized diagnostic needs in medical imaging tasks. By applying the personalized federated learning method of the present invention, consistent feature-level alignment and adaptive optimization of local models can be achieved without directly exchanging raw medical data, thereby effectively improving the classification performance of each client in specific medical imaging tasks. BRIEF DESCRIPTION OF THE DRAWINGS
[0025] Figure 1Schematic diagram of the process of the personalized federated learning method used in the present invention; Figure 2 This is a schematic diagram of the framework of the training phase of the personalized federated learning method used in the present invention; Figure 3 This is a schematic diagram of the framework of the reasoning phase of the personalized federated learning method used in the present invention; Figure 4 This is a comparison chart of the test accuracy of the personalized federated learning algorithm of the present invention and other federated learning algorithms using CIFAR100 data in Example 2; Figure 5 This is Example 2, a comparison chart of the test accuracy of the personalized federated learning algorithm of the present invention and other federated learning algorithms under Digit5 data. DETAILED DESCRIPTION
[0026] Symbol definition: This invention selects clients participate in the aggregation, each client Its local training set is , where | | indicates the number of datasets it has. Both local training sets and local test sets are image datasets. Global dataset Represents the collection of all client datasets.
[0027] Example 1 The present invention provides a personalized federated learning method for classifier asynchronous optimization and prototype-aware reasoning, such as Figure 1 、 2 , 3, the method includes: Training phase: S0, server-side initializes shared classifier and global clustering prototypes , initialize the shared classifier and global clustering prototypes Broadcast to each client participating in the training; S1. Asynchronous update: The client receives the shared classifier and global clustering prototype from the server. The client first freezes the personalized classifier, uses the global clustering prototype to adaptively align features, guides the update of the personalized feature extractor and shared classifier, and unfreezes the personalized classifier after the update. Subsequently, the client freezes the updated personalized feature extractor and shared classifier again, updates the personalized classifier, and unfreezes the personalized feature extractor and shared classifier after the update is complete. Step S1 uses the global clustering prototype to adaptively perform feature alignment to guide the update of the personalized feature extractor and the shared classifier, specifically: S11. The client calculates the entropy value based on the label distribution of the local training set to obtain the adaptive alignment weight for feature alignment. : Specifically, in traditional feature alignment methods, all clients use the same alignment weight. That is to say, each client will align the features with the same effort, without taking into account the differences in data distribution among the clients. When the data differences between clients are very serious, forcibly using the same alignment strength will cause the following problems: ① The gradient of the alignment item becomes larger: the system will "try hard" to align features with large distribution differences. As a result, during the calculation process, the gradient during back propagation will become particularly large; ② Interference with the main task learning: the alignment operation is over-enhanced, which will hinder the learning process of the main task (such as classification). Therefore, the present invention designs an adaptive feature alignment method to make the alignment weight of clients with uneven data distribution smaller. Here we use an entropy-based method to calculate the alignment weight. First, the label distribution of each client is counted, and then the entropy of the label distribution is calculated using the following formula: (1) In formula (1), Represents the random variable Y taking the tth category in the label space The probability of C represents the total number of categories in the label space. represents the t-th category in the label space, represent The logarithm of The entropy of the label distribution of client k. A larger value indicates a more uneven label distribution, which is used to evaluate the degree of heterogeneity of data distribution. Next, the obtained entropy value is used to calculate the adaptive alignment weight of client k: (2) In formula (2), is the adaptive alignment weight of client k, is the entropy of the label distribution of client k, is the scaling factor, the scaling factor The role of is to normalize the alignment term to prevent the entropy value itself from being too large or too small, which in turn affects the stability of training. It can also adjust the influence ratio of the alignment term relative to the main task, thereby setting a reasonable balance between the main task and the alignment task.
[0028] S12, based on adaptive alignment weight , update the personalized feature extractor and shared classifier as follows: For updating the two classification heads, an asynchronous update strategy is used, that is, updating the personalized feature extractor and shared classifier first, and then updating the personalized classifier, to prevent synchronous updates from interfering with the training of the two. First, the personalized classifier is frozen, and then the personalized feature extractor and shared classifier are updated. Specifically, for updating the personalized feature extractor, the global clustering prototype obtained from the server is combined with contrastive learning to guide it to produce semantically consistent and class-separable features. That is, features of the same class are close to each other, and features of different classes are far away from each other. The loss of feature alignment is calculated as follows: (3) In formula (3), is the loss function of the personalized feature extractor, The feature representation of the i-th sample extracted by the personalized feature extractor of the k-th client before updating, represent The cosine similarity between and the category prototype c, Represents the category The set of global cluster prototypes is the positive sample set, Representatives do not belong to the category The set of global clustering prototypes is the set of negative samples. By calculating this loss, the personalized feature extractors of each client produce semantically consistent feature representations while maintaining separability between categories. Contrastive learning is used to guide the feature alignment of the client's personalized feature extractors, thereby addressing domain shift.
[0029] Next, for the update of the shared classifier, the traditional cross entropy loss is used: (4) In formula (4), is the cross entropy loss of the shared classifier, | | represents the local training set of client K The number of samples, is the label value of the i-th sample, Represents the prediction value of the shared classifier for the i-th sample, and the shared classifier is updated by minimizing the loss; The total loss of the personalized feature extractor and shared classifier update is: (5) In formula (5), is the total loss of the update, is the cross entropy loss of the shared classifier, represents the adaptive weight for feature alignment, is the loss function of the personalized feature extractor, which improves the performance of the personalized feature extractor and the shared classifier by minimizing the updated total loss, and obtains the updated personalized feature extractor and shared classifier; Step S1 of updating the personalized classifier is as follows: Use cross entropy loss to update it, and its loss function is: (6) (7) in, is the total loss of the personalized classifier update, represents the cross entropy loss of the personalized classifier, | | represents the local training set of client K The number of samples, is the label value of the i-th sample, Represents the predicted output of the personalized classifier of the i-th sample; by minimizing To improve the classification performance of the personalized classifier and obtain an updated personalized classifier.
[0030] S2. Local prototype clustering: Use the updated personalized feature extractor to extract features from the local training set, and perform local prototype clustering on the extracted features to obtain a local cluster prototype set. , and then perform weighted average calculation on the prototypes in the set to obtain the local unbiased prototype ; The details are as follows: Local prototype clustering: Current research on local prototypes has all used the local average prototype method, but this prototype tends to favor a dominant feature and ignore other important features. Therefore, we consider using the FINCH clustering method to obtain local cluster prototypes. Then, we average the local cluster prototypes belonging to the same category to obtain a local unbiased prototype. This prototype does not favor a dominant feature and better reflects the average characteristics of a domain, thus avoiding information loss.
[0031] First, client k performs FINCH clustering on the features extracted from the personalized feature extractor. The clustering process is as follows: (8) In formula (8), represents the set of local cluster prototypes of client k belonging to category m, Represents the jth feature prototype of client k belonging to category m, represents the number of cluster prototypes of client k belonging to category m, Represents the input sample The features obtained by the updated personalized feature extractor, represents the i-th input sample and its label, The local training set represents the k-th client belonging to category m. After obtaining the feature prototype of each category, the feature prototypes of each category are averaged to obtain the local unbiased prototype. The calculation process is as follows: (9) In formula (9), represents the local unbiased prototype of client k belonging to category m, Represents the jth feature prototype of client k belonging to category m, Represents the number of cluster prototypes for which client k belongs to category m.
[0032] S3, the client uses the shared classifier updated in step S1 and the local unbiased prototype obtained in step S2 Send to server; S4, shared classifier aggregation: The server aggregates the received shared classifiers using the average aggregation method to obtain the aggregated shared classifiers ; Specifically, the server aggregates the local shared classifier using the following formula: (10) In formula (10), represents the aggregated global shared classifier, Represents the number of clients participating in the aggregation, is the local shared classifier of the k-th client.
[0033] S5, global prototype clustering: the server will receive the local unbiased prototype Perform global prototype clustering to obtain the global cluster prototype ; In order to better deal with domain shift, the global prototype needs to contain more domain diversity knowledge. Since the ordinary global average prototype cannot describe different domain information and is biased towards the underlying dominant domain, the present invention uses global clustering prototypes to supplement the rich domain diversity knowledge. In order to obtain the global clustering prototype, the method of the present invention uses FINCH to cluster the uploaded local unbiased prototypes. The clustering process is as follows: (11) In formula (11), represents the set of global cluster prototypes of category m, represents the t-th global cluster prototype of category m, represents the number of global cluster prototypes for category m, represents the set of local unbiased prototypes for each client class m, A local unbiased prototype representing class m for client k.
[0034] S6. The server uses the aggregated shared classifier obtained in step S4 and the global clustering prototype obtained in step S5 Rebroadcast to all participating clients and repeat steps S1-S6 until the preset rounds are reached or the model converges; Adaptive reasoning stage based on prototype perception: S7. After training is completed, the corresponding features extracted from the local test set by the personalized feature extractor are input into the shared classifier and the personalized classifier to obtain personalized prediction output and shared prediction output. Based on prototype perception technology, the client adaptively calculates the output weights of the two classifiers and performs weighted fusion on the prediction outputs of the two classifiers to obtain the final prediction result, and then obtains the final test accuracy, which is specifically: To adaptively integrate the prediction outputs of the personalized and shared heads, thereby improving the model's adaptability to data distribution, this paper uses the cosine similarity between features and prototypes to approximate the confidence of the two classifiers. If the current feature is more similar to the global prototype, the shared classifier is more trusted. Conversely, if the feature is more similar to the local unbiased prototype, the personalized classifier is more trusted. The specific process is as follows: S71. After the client is trained, first, the local test set is fed into the personalized feature extractor to extract corresponding features, and the features are fed into the personalized classifier and the shared classifier to obtain personalized prediction output and shared prediction output respectively. S72, calculate the features extracted by the personalized feature extractor and the global prototype and locally unbiased prototypes The cosine similarity of Calculating the similarity between features and prototypes: To measure the confidence between the two, we first need to calculate the similarity between the features generated by the personalized feature extractor and the global prototype, as well as the cosine similarity between the features generated by the personalized feature extractor and the local prototype. The global prototype here refers to the average of the global cluster prototypes of the corresponding category, and the local prototype refers to the local unbiased prototype. The process of calculating cosine similarity is as follows: (12) (13) (14) In formula (12), (13), (14), For the client Features generated by the personalized feature extractor during the inference phase and corresponding categories Locally unbiased prototype of The cosine similarity between For the client Features generated by the personalized feature extractor during the inference phase and corresponding categories The global prototype The cosine similarity between yes The standardized measure of yes The standardized measure of yes The standardized measure of represents the number of global cluster prototypes for category m, Represents the tth global cluster prototype of category m.
[0035] S73, using the softmax function to normalize the two cosine similarities calculated in step S72 to obtain weights of the prediction outputs of the shared classifier and the personalized classifier, and using the weights to fuse the two prediction outputs obtained in step S71 to obtain the final inference prediction output; Specifically, the prototype perception weight is calculated: Based on the calculated similarity, the weight between the prediction output of the shared classifier and the prediction output of the personalized classifier is calculated using the following formula: (15) In formula (15), and represent the weight of the prediction output of the shared classifier and the weight of the prediction output of the personalized classifier for client k, respectively, and + =1, For the client Features generated by the personalized feature extractor and corresponding categories Locally unbiased prototype of The cosine similarity between For the client Features generated by the personalized feature extractor and corresponding categories The global prototype The cosine similarity between .
[0036] Finally, based on the weights of the prediction outputs of the personalized classification head and the shared classification head, the final inference prediction output is: (16) in, is the inference prediction output, and represent the prediction output weights of the shared classifier and the prediction output weights of the personalized classifier for client k, respectively. is the prediction output of the personalized classifier for client k, is the prediction output of the shared classifier.
[0037] S74. Compare the final inference prediction output obtained in step S73 with the actual output result to obtain the final test accuracy.
[0038] Figure 2 Step 1 to step 6 refer to S1 to S6 of the present invention.
[0039] The present invention also provides an application of a personalized federated learning method of asynchronous optimization of classifiers and prototype-aware reasoning in medical image analysis, which is used for personalized model optimization tasks in medical image analysis.
[0040] Example 2 Next, we will conduct experimental verification. First, we will give a basic introduction to the experiment: (1) Dataset Introduction: CIFAR100: This dataset consists of 60,000 color images, each of size 32×32 pixels, organized into 100 different categories, each containing 6,000 images. Digit5: It contains handwritten digit image datasets from five different domains: MNIST, MNIST-M, SVHN, SYN, and USPS. The data from each domain has different style, background, and other characteristics. (2) Task: image classification task; (3) Model: CNN, consisting of three convolution-pooling layers and two fully connected layers; (4) Label offset setting: The CIFAR100 dataset is used for training and testing. The data partitioning method adopts practical non-IID partitioning (α=0.1), and the distribution of the training set and the test set is the same; (5) Domain-offset setting: The Digit5 dataset is used for training and testing. During training, there are five clients, each with data from one domain. The data of each client belongs to a different domain. During testing, the test set of each client comes from the data of the other client's domain.
[0041] (6) Benchmark algorithms: FedAvg, FedProx, FedPer, FedRep, FedPAC, FedProto, FedAMP, FedAPEN, FedKD, FML; (7) Hyperparameter settings: The learning rate of FedAvg and FedProx is set to 0.1, and the learning rate of FedPer, FedRep, FedPAC, FedProto, FedAMP, FedAPEN, FedKD, FML, and our algorithm is set to 0.05. The number of local iterations for all algorithms is 3, and the batch size is 32; To ensure fairness, all benchmark algorithms and this algorithm use the same network architecture, equipment, and hyperparameter settings. Figure 3 and 4 The experimental results are analyzed in detail.
[0042] For attached Figure 4 This section compares the performance of the proposed personalized federated learning algorithm with other baseline methods on the CIFAR-100 dataset, using a heterogeneity coefficient of α = 0.1. As can be seen from the figure, traditional federated learning methods, such as FedAvg and FedProx, perform significantly worse in heterogeneous data scenarios than personalized federated learning methods, including FedPer, FedRep, FedProto, FedAPEN, FedPAC, and the method of the present invention. This discrepancy is due to the fact that traditional federated learning methods primarily optimize the model's global generalization performance and lack attention to client-side personalized needs. Personalized federated learning methods, on the other hand, enhance performance in heterogeneous data conditions through local adaptability. It is particularly noteworthy that FedPer and FedRep exhibit significant performance degradation compared to other personalized methods. This is primarily due to their use of a shared feature extractor structure, which forces the model to generate more domain-invariant features. While this feature representation improves cross-domain generalization, it can mislead the classifier in personalized classification tasks, leading to performance degradation. In contrast, the method proposed in this paper performs the best among all the compared methods, verifying the effectiveness and superiority of the designed asynchronous classifier update strategy and prototype-aware adaptive reasoning strategy in label heterogeneous environments.
[0043] For attached Figure 5 The Digit5 dataset was used to evaluate the adaptability of our method in scenarios with domain shift. Experimental results show that compared to other personalization methods, our method demonstrates superior performance when migrating between different source and target domains, further demonstrating the effectiveness of the proposed bilateral prototype clustering strategy in feature alignment and domain generalization.
[0044] In summary, the above experimental results fully demonstrate that the method of the present invention has significant advantages in dealing with personalized federated learning scenarios with label shift and domain shift, and each key module plays a vital role in improving the overall performance.
Claims
1. A personalized federated learning method based on asynchronous optimization of classifiers and prototype-aware reasoning, characterized by: The method comprises: Training phase: S0, server-side initializes shared classifier and global clustering prototypes , initialize the shared classifier and global clustering prototypes Broadcast to each client participating in the training; S1. Asynchronous update: The client receives the shared classifier and global clustering prototype from the server. The client first freezes the personalized classifier, uses the global clustering prototype to adaptively align features, and guides the update of the personalized feature extractor and shared classifier. After the update, the personalized classifier is unfrozen. The client then freezes the updated personalized feature extractor and shared classifier, updates the personalized classifier, and unfreezes them after the update. S2. Use the updated personalized feature extractor to extract features from the local training set, and perform local prototype clustering on the extracted features to obtain a local cluster prototype set. , and then perform weighted average calculation on the prototypes in the set to obtain the local unbiased prototype ; S3, the client uses the shared classifier updated in step S1 and the local unbiased prototype obtained in step S2 Send to server; S4. The server aggregates the received shared classifiers using the average aggregation method to obtain the aggregated shared classifiers. ; S5. The server receives the local unbiased prototype Perform global prototype clustering to obtain the global cluster prototype ; S6. The server uses the aggregated shared classifier obtained in step S4 and the global clustering prototype obtained in step S5 Rebroadcast to all participating clients and repeat steps S1-S6 until the preset rounds are reached or the model converges; Adaptive reasoning stage based on prototype perception: S7. After training is completed, the corresponding features of the local test set extracted by the personalized feature extractor are input into the shared classifier and the personalized classifier to obtain personalized prediction output and shared prediction output. Based on prototype perception technology, the client adaptively calculates the output weights of the two classifiers and performs weighted fusion on the prediction outputs of the two classifiers to obtain the final prediction result, and then obtains the final test accuracy.
2. The personalized federated learning method for asynchronous classifier optimization and prototype-aware reasoning according to claim 1 is characterized in that: Step S1 uses the global clustering prototype to adaptively perform feature alignment to guide the update of the personalized feature extractor and the shared classifier, specifically: S11. The client calculates the entropy value based on the label distribution of the local training set to obtain the adaptive alignment weight for feature alignment. ; Statistics on the label distribution of each client, and use formula (1) to calculate the entropy of label distribution: (1) In formula (1), represents the entropy of the label distribution of client k, Represents the random variable Y taking the tth category in the label space The probability of represents the t-th category in the label space, C represents the total number of categories in the label space, represent The logarithm of , the larger its value, the more uneven the label distribution, which is used to evaluate the degree of heterogeneity of data distribution; Use the obtained entropy value to calculate the adaptive alignment weight of client k : (2) In formula (2), is the adaptive alignment weight of client k, is the entropy of the label distribution of client k, is the scaling factor; S12, based on adaptive alignment weight , update the personalized feature extractor and shared classifier: For updating the personalized feature extractor, we use the global clustering prototype previously obtained from the server in conjunction with contrastive learning to guide it to generate semantically consistent and class-separable features. That is, features of the same class are close to each other, and features of different classes are far away from each other. The loss of feature alignment is calculated as follows: (3) In formula (3), is the loss function of the personalized feature extractor, The feature representation of the i-th sample extracted by the personalized feature extractor of the k-th client before updating, represent The cosine similarity between and category prototype c, Represents the category The set of global cluster prototypes is the positive sample set, Representatives do not belong to the category The set of global clustering prototypes is the negative sample set. By calculating this loss, the personalized feature extractor of each client can produce semantically consistent feature representations while maintaining the separability between categories. For the update of the shared classifier, the traditional cross entropy loss is used: (4) In formula (4), is the cross entropy loss of the shared classifier, | | represents the local training set of client K The number of samples, is the label value of the i-th sample, Represents the prediction value of the shared classifier for the i-th sample, and the shared classifier is updated by minimizing the loss; The total loss of the personalized feature extractor and shared classifier update is: (5) In formula (5), is the total loss of the update, is the cross entropy loss of the shared classifier, represents the adaptive weight for feature alignment, is the loss function of the personalized feature extractor, which is minimized by , to improve the performance of the personalized feature extractor and the shared classifier, and obtain the updated personalized feature extractor and the shared classifier; Step S1 of updating the personalized classifier is as follows: Use cross entropy loss to update it, and its loss function is: (6) (7) in, is the total loss of personalized classifier update, represents the cross entropy loss of the personalized classifier, | | represents the local training set of client K The number of samples, is the label value of the i-th sample, Represents the predicted output of the personalized classifier of the i-th sample; by minimizing To improve the classification performance of the personalized classifier and obtain an updated personalized classifier.
3. The personalized federated learning method for asynchronous classifier optimization and prototype-aware reasoning according to claim 1 is characterized in that: In step S2, the extracted features are clustered into local prototypes to obtain a local cluster prototype set. , specifically: Client k performs FINCH clustering on the features extracted from the personalized feature extractor. The clustering process is as follows: (8) In formula (8), represents the set of local cluster prototypes of client k belonging to category m, Represents the jth feature prototype of client k belonging to category m, represents the number of cluster prototypes of client k belonging to category m, Represents the input sample The features obtained by the updated personalized feature extractor, represents the i-th input sample and its label, represents the local training set of the kth client belonging to category m; The weighted average calculation of the prototypes in the set described in step S2 is performed to obtain the local unbiased prototype , specifically: After obtaining the feature prototype of each category, the feature prototypes of each category are averaged to obtain the local unbiased prototype. The calculation process is as follows: (9) In formula (9), represents the local unbiased prototype of client k belonging to category m, Represents the jth feature prototype of client k belonging to category m, Represents the number of cluster prototypes for which client k belongs to category m.
4. The personalized federated learning method for asynchronous classifier optimization and prototype-aware reasoning according to claim 1 is characterized in that: The step S4 is specifically as follows: The server aggregates the local shared classifier using the following formula: (10) In formula (10), represents the shared classifier after aggregation, Represents the number of clients participating in the aggregation, is the local shared classifier of the k-th client.
5. The personalized federated learning method for asynchronous classifier optimization and prototype-aware reasoning according to claim 1 is characterized in that: In step S5, the server receives the local unbiased prototype Perform global prototype clustering to obtain the global cluster prototype , specifically: Use FINCH to upload local unbiased prototypes Clustering is performed, and the clustering process is as follows: (11) In formula (11), represents the set of global cluster prototypes of category m, represents the t-th global cluster prototype of category m, represents the number of global cluster prototypes for category m, represents the local unbiased prototype of category m for client k, A set of local unbiased prototypes representing each client class m.
6. The personalized federated learning method for asynchronous classifier optimization and prototype-aware reasoning according to claim 1 is characterized in that: The step S7 is specifically as follows: S71. After training is completed, first, the local test set is fed into the personalized feature extractor to extract corresponding features, and the features are fed into the personalized classifier and the shared classifier to obtain personalized prediction output and shared prediction output respectively; S72, calculate the features extracted by the personalized feature extractor and the global prototype and locally unbiased prototypes The cosine similarity of The prototype of global clustering The average calculation is obtained; S73, using the softmax function to normalize the two cosine similarities calculated in step S72 to obtain weights of the prediction outputs of the shared classifier and the personalized classifier, and using the weights to fuse the two prediction outputs obtained in step S71 to obtain the final inference prediction output; S74. Compare the final inference prediction output obtained in step S73 with the actual output result to obtain the final test accuracy.
7. The personalized federated learning method for asynchronous classifier optimization and prototype-aware reasoning according to claim 6 is characterized in that: Step S72 calculates the features extracted by the personalized feature extractor and the features extracted by the global prototype and locally unbiased prototypes The cosine similarity is: (12) (13) (14) In formula (12), (13), and (14), For the client Features generated by the personalized feature extractor during the inference phase and corresponding categories Locally unbiased prototype of The cosine similarity between For the client Features generated by the personalized feature extractor during the inference phase and corresponding categories The global prototype The cosine similarity between yes The standardized measure of yes The standardized measure of yes The standardized measure of represents the number of global cluster prototypes for category m, Represents the t-th global cluster prototype of category m.
8. The personalized federated learning method for asynchronous classifier optimization and prototype-aware reasoning according to claim 6, characterized in that: The step S73 is specifically as follows: The weight between the prediction output of the shared classifier and the prediction output of the personalized classifier is calculated by formula (15): (15) In formula (15), and represent the weight of the prediction output of the shared classifier and the weight of the prediction output of the personalized classifier for client k, respectively, and + =1, For the client Features generated by the personalized feature extractor and corresponding categories Locally unbiased prototype of The cosine similarity between For the client Features generated by the personalized feature extractor and corresponding categories The global prototype The cosine similarity between According to the weights of the prediction outputs of the personalized classifier and the shared classifier, the final inference prediction output is: (16) in, is the inference prediction output, and represent the prediction output weights of the shared classifier and the prediction output weights of the personalized classifier for client k, respectively. is the prediction output of the personalized classifier for client k, is the prediction output of the shared classifier.
9. Application of the personalized federated learning method of asynchronous classifier optimization and prototype-aware reasoning as described in any one of claims 1 to 8 in medical image analysis, for personalized model optimization tasks in medical image analysis.
Citation Information
Patent Citations
Single-client multi-domain heterogeneous federated learning system and method based on manifold learning
CN118734995A
Special equipment for fitting plastic mold by bench worker
CN119658639A
Federal continuous learning method based on improved unet network
CN118690832A
Federal learning method based on flexible combination of feature extractor and classifier
CN119005302A
Client demand-oriented industrial Internet of Things personalized federal learning method
CN119168090A
Cited By
Personalized federal learning method applied to privacy calculation
CN115660107A