Federal domain generalization-oriented medical image classification method and system
Through the combination of image style transfer and JS divergence consistency loss, the problem of insufficient model overfitting and cross-domain generalization capabilities in federated learning is solved, and more efficient model generalization and accuracy improvement in medical image classification tasks are achieved.
Patent Information
- Application Number
- CN202510638600.9
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-05-19
- Publication Date
- 2025-08-15
AI Technical Summary
When facing the problem of non-independent and homogeneous data heterogeneity, existing federated learning methods are difficult to improve the cross-domain generalization capabilities of models while protecting data privacy. Especially in medical image classification tasks, traditional data-level, model-level and server-level solutions are difficult to effectively adapt to the distribution characteristics of unknown target test domains.
The visible data domain is expanded through image style migration technology, and a multi-style training data set is generated using global style sets, and JS divergence consistency loss is introduced as a regular term in the client local model. Combined with the global data set, the generalization ability of the client local model is evaluated, and the aggregation weight is dynamically adjusted to optimize the model aggregation effect.
It effectively alleviates the problem of model overfitting, improves the generalization performance of the model in unknown target domains, and improves the accuracy and aggregation effect of medical image classification.
Smart Images

Figure CN120495806A_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the field of image classification technology, and in particular to a medical image classification method and system for generalization in a federated domain. Background Art
[0002] In recent years, artificial intelligence (AI) technology has become deeply integrated into human society, becoming a core engine driving the development of modern civilization. As AI applications continue to expand, the dimensions and scale of data collection, a key resource for this intelligent revolution, continue to expand. The resulting issues of data security and privacy protection have become a global concern. Under the traditional centralized machine learning paradigm, all training data must be aggregated and processed on a central server. This model not only presents the risk of a single point of failure but is also highly susceptible to large-scale data leaks. Particularly in sensitive sectors such as healthcare and financial transactions, cross-institutional data sharing faces strict legal and regulatory restrictions and technical barriers due to industry regulations and the need to protect data privacy. Against this backdrop, federated learning, an emerging distributed machine learning paradigm, has emerged. The core concept of federated learning is to enable multiple clients (such as hospitals and banks) to collaboratively train a global model without sharing local data, thereby achieving model optimization while protecting data privacy.
[0003] However, federated learning faces numerous challenges in its practical application, such as data heterogeneity, particularly the non-independent and identically distributed (Non-IID) problem. While a large body of research has made significant progress in addressing the Non-IID problem in federated learning, most of this work relies on an idealized assumption: that the test data and training data come from the same distribution. This results in model optimization tailored to the known data domain, neglecting more realistic application scenarios where the training data exhibits significant domain shift and the target test domain is completely invisible during federated training. In these situations, global models trained on distributed multi-source domain data often struggle to adapt to the distributional characteristics of the unknown target test domain. Therefore, exploring global model training methods with cross-domain generalization capabilities within the federated learning framework (i.e., federated domain generalization methods) is crucial for promoting the application of federated learning in real-world scenarios.
[0004] Existing federated domain generalization approaches primarily improve model generalization capabilities at the data, model, and server levels. Data-level approaches enrich data diversity through cross-client information exchange. However, due to limited variations in augmented images or inappropriate local training methods, these approaches often lead to model overfitting to the training data domain, resulting in limited improvement in generalization. Research at the model level focuses on two main directions: first, network architecture innovation, through the design of novel structures or the adoption of cutting-edge architectures to extract more robust domain-invariant features; second, improved training strategies, such as utilizing adversarial learning to align source domain features or introducing regularization to prevent model overfitting. These approaches have achieved some success, but unlike data-level approaches, they employ traditional server-side average aggregation or aggregation based on data volume proportion, ignoring the differences in generalization capabilities of local client models and resulting in suboptimal aggregation results. While server-level approaches attempt to assign aggregation weights based on model generalization capabilities, their evaluation metrics often fail to accurately measure the model's true generalization performance in unknown target domains. Summary of the Invention
[0005] The purpose of the present invention is to overcome the shortcomings of the existing technology and provide a medical image classification method for federated domain generalization, which can expand the visible data domain of model training, prevent the model from overfitting the client's private data domain, and improve the model aggregation effect and model generalization performance based on the client's local model generalization performance evaluation results.
[0006] To achieve the above object, the present invention is implemented by adopting the following technical solutions:
[0007] In one aspect, the present invention provides a medical image classification method for federated domain generalization, comprising:
[0008] Obtaining medical images to be classified;
[0009] Inputting the medical image to be classified into a pre-trained image classification model and outputting a medical image classification result;
[0010] The training of the image classification model includes:
[0011] Obtain historical medical image styles for each client;
[0012] Generate a global style set based on the historical medical image styles of each client, and use the global style set to perform image style transfer on the original training dataset to generate a multi-style training dataset; the multi-style training dataset includes a private dataset and a global dataset;
[0013] The image classification model to be trained is distributed to each client as the client-side local model to be trained. The private dataset is input into the client-side local model to obtain a trained client-side local model. JS divergence consistency loss is added to the objective function of the client-side local model as a regularization term.
[0014] Receive the trained local models of each client, calculate the generalization ability score of each client local model based on the global data set, calculate the aggregation weight of the client local model based on the generalization ability score, and perform weighted aggregation on the client local models according to the aggregation weight to obtain a trained image classification model.
[0015] Optionally, use the global style set to perform image style transfer on the original training dataset to generate a multi-style training dataset, including:
[0016] Obtaining an original training dataset; the original training dataset includes an original private dataset and an original global dataset;
[0017] Randomly select the historical medical image style of other clients in the global style set as the privacy transfer style, perform style transfer on the original private dataset according to the privacy transfer style, and obtain a privacy auxiliary dataset. Then, obtain the private dataset based on the local private dataset and the privacy auxiliary dataset.
[0018] The historical medical image styles of the client in the global style set are fused to obtain a fused global style; the fused global style is used as the global migration style, and the original global dataset is style migrated according to the global migration style to obtain a global auxiliary dataset; and a global dataset is obtained based on the original global dataset and the global auxiliary dataset.
[0019] Optionally, fusing the historical medical image styles of the client in the global style set to obtain a fused global style includes:
[0020] ;
[0021] ;
[0022] ;
[0023] in, represents the historical medical image style of the kth client; 、 denote the mean and standard deviation of the k-th client image feature respectively; Represents the mathematical expectation of the original private data; represents an image sample; represents the original private dataset of the kth client; is the result after data encoding; Indicates the calculation of the mean, Indicates the calculation of variance; Represents a global style set; Indicates the number of clients; Indicates the integration of global style; represents the style fusion weight of the k-th client.
[0024] Optionally, the objective function of the client local model is expressed as:
[0025] ;
[0026] in, represents the objective function of the local model of the kth client in the tth iteration; represents the k-th client privacy dataset; represents the k-th client local model in the t-th iteration; Represents the mathematical expectation on the private dataset; represents the cross entropy loss function; Indicates calculation of JS divergence; Indicates the strength of JS consistency restrictions; represents the predicted distribution of the original private data of the kth client in the tth iteration; 、 represents the predicted distribution of the k-th client's privacy auxiliary data in the t-th iteration; represents the original private data of the kth client; 、 represents the k-th client privacy auxiliary data; Indicates the label corresponding to the original private data of the kth client.
[0027] Optionally, the updating process of the client local model is expressed as:
[0028] ;
[0029] in, Indicates the k-th client local model in the t+1th iteration; Indicates the update model; represents the k-th client local model in the t-th iteration; Indicates the client local learning rate; Indicates the gradient is obtained; represents the local objective function of the kth client; represents the k-th client private dataset.
[0030] Optionally, the generalization ability score of each client local model is expressed as:
[0031] ;
[0032] in, represents the generalization ability score of the k-th client local model in the t-th iteration; represents the predicted distribution of the original global data in the tth iteration; 、 represents the predicted distribution of the global auxiliary data in the tth iteration; Represents original global data; 、 Represents global auxiliary data; represents the k-th client local model in the t-th iteration; Indicates the number of original global data; represents the KL divergence.
[0033] Optionally, calculate the aggregate weights of the image classification model based on the generalization ability score, including:
[0034] ;
[0035] ;
[0036] in, represents the initial aggregate weight of the k-th client in the t-th iteration; represents the amount of local private data of the kth client; represents the influencing factor of generalization ability score; represents the generalization ability score of the k-th client local model in the t-th iteration; represents the final aggregate weight of the k-th client in the t-th iteration; Indicates the number of clients.
[0037] Optionally, the trained image classification model is expressed as:
[0038] ;
[0039] in, Represents the t+1th round of iterative image classification model; represents the final aggregate weight of the k-th client in the t-th iteration; represents the k-th client local model in the t-th iteration; Indicates the number of clients.
[0040] In a second aspect, the present invention provides a medical image classification system generalized for federated domains, comprising:
[0041] An image acquisition module is used to: acquire medical images to be classified;
[0042] An image classification module is used to: input the medical image to be classified into a pre-trained image classification model and output a medical image classification result;
[0043] A model training module is used for: wherein the training of the image classification model includes:
[0044] Obtain historical medical image styles for each client;
[0045] Generate a global style set based on the historical medical image styles of each client, and use the global style set to perform image style transfer on the original training dataset to generate a multi-style training dataset; the multi-style training dataset includes a private dataset and a global dataset;
[0046] The image classification model to be trained is distributed to each client as the client-side local model to be trained. The private dataset is input into the client-side local model to obtain a trained client-side local model. JS divergence consistency loss is added to the objective function of the client-side local model as a regularization term.
[0047] Receive the trained local models of each client, calculate the generalization ability score of each client local model based on the global data set, calculate the aggregation weight of the client local model based on the generalization ability score, and perform weighted aggregation on the client local models according to the aggregation weight to obtain a trained image classification model.
[0048] In a third aspect, the present invention provides a computer-readable storage medium having a computer program stored thereon, which, when executed by a processor, implements the steps of the medical image classification method for federated domain generalization as described in the first aspect.
[0049] Compared with the prior art, the present invention has the following beneficial effects:
[0050] The present invention uses image style transfer technology to expand the visible data domain of model training and prevent the model from overfitting the client's private data domain through local regularization. On the other hand, it uses multi-style versions of global data to simulate an unknown domain environment to test the generalization performance of the client's local model. Based on the generalization performance evaluation results, it dynamically adjusts the contribution ratio of the client's local model to the image classification model, further improving the aggregation effect and model generalization performance. BRIEF DESCRIPTION OF THE DRAWINGS
[0051] Figure 1 FIG2 is a flow chart of a generalized medical image classification method for federated domains according to an embodiment of the present invention;
[0052] Figure 2 FIG2 is a schematic diagram of the process of image style transfer in one embodiment of the present invention;
[0053] Figure 3 FIG2 is a schematic diagram showing the result of re-stylization of a PACS dataset in one embodiment of the present invention;
[0054] Figure 4 FIG2 is a schematic diagram showing the result of re-stylizing the Office-Home dataset in one embodiment of the present invention;
[0055] Figure 5 FIG2 is a schematic diagram showing the result of re-stylizing the PACS dataset as a global dataset in one embodiment of the present invention;
[0056] Figure 6 FIG2 is a schematic diagram showing the result of re-stylizing the Office-Home dataset as a global dataset in one embodiment of the present invention;
[0057] Figure 7 The figure shows the test accuracy change curves of various methods in one embodiment of the present invention using the Art domain of the PACS dataset as the unknown target domain;
[0058] Figure 8 The figure shows the test accuracy change curves of various methods in one embodiment of the present invention using the Cartoon domain of the PACS dataset as an unknown target domain;
[0059] Figure 9 The figure shows the test accuracy change curve of each method in an embodiment of the present invention using the Photo domain of the PACS dataset as the unknown target domain;
[0060] Figure 10 The figure shows the test accuracy change curves of various methods in one embodiment of the present invention using the Sketch domain of the PACS dataset as an unknown target domain;
[0061] Figure 11 The figure shows the test accuracy change curve of each method in one embodiment of the present invention using the Art domain of the Office-Home dataset as the unknown target domain;
[0062] Figure 12 The figure shows the test accuracy change curve of each method in an embodiment of the present invention using the Clipart domain of the Office-Home dataset as an unknown target domain;
[0063] Figure 13 The figure shows the test accuracy change curve of each method in an embodiment of the present invention using the Product domain of the Office-Home dataset as the unknown target domain;
[0064] Figure 14Shown is the test accuracy change curve of each method in an embodiment of the present invention using the Real_World domain of the Office-Home dataset as the unknown target domain. DETAILED DESCRIPTION
[0065] The technical solution of the present invention is described in detail below through the accompanying drawings and specific embodiments. It should be understood that the embodiments of the present invention and the specific features in the embodiments are detailed descriptions of the technical solution of the present invention, rather than limitations on the technical solution of the present invention. In the absence of conflict, the embodiments of the present invention and the technical features in the embodiments can be combined with each other.
[0066] The term "and / or" simply describes a relationship between related objects, indicating that three possible relationships exist. For example, "A and / or B" can mean: A exists alone, A and B exist simultaneously, or B exists alone. Additionally, the character " / " generally indicates an "or" relationship between the related objects.
[0067] Example 1
[0068] like Figure 1 As shown, this embodiment introduces a medical image classification method for federated domain generalization, including the following steps:
[0069] Step 1: Obtain the medical image to be classified.
[0070] Step 2: Input the medical image to be classified into the pre-trained image classification model and output the medical image classification results, specifically:
[0071] The image classification model is processed at the server level and, after initialization, sent to the client at the client level for local training. The trained local model is then uploaded to the server. The server optimizes the aggregation weights of the client's local model through generalization evaluation and performs aggregation accordingly. This process is repeated multiple times to obtain the final image classification model.
[0072] Among them, at the client layer: This layer contains clients, each client Has a size of The original privacy dataset ,in Represents the original privacy data, Original private data The corresponding labels, and assuming that the original private dataset follows the distribution From the perspective of being closer to actual application scenarios, the data domains of each client are independent of each other, and there is a domain offset, that is, , The existence of domain shift will cause the model trained on the distributed client to be only applicable to its local data domain, while the global model obtained through aggregation will be ineffective when applied to an unknown target domain (i.e., the data distribution is not visible during training, i.e., the test data domain distribution is not visible). , ( ) cannot maintain good discrimination. To address this problem, this embodiment proposes that before the training begins, each client Compute historical medical image styles in its local data domain , and upload it to the server. Then, the client will download the global style collection collected by the server Perform image style transfer to obtain re-stylized privacy-assisted data 、 , to conduct auxiliary training. When the global training of the round of iteration begins, the client First, use the global model to be trained sent by the server As its initialization local model, then based on the original private data and restyled privacy-assisted data 、 to carry out Round of local training. Finally, the client local model is obtained and upload it to the server for aggregation.
[0073] Server layer: This layer mainly consists of a server with powerful computing and storage capabilities, which covers the client layer clients, and deployed a Unlabeled raw global dataset This layer has the following three functions: 1) Style collection and broadcasting. Before training begins, collect the styles of each client. Historical medical image styles , forming a global style set , and broadcast it to each client. 2) Aggregation weight optimization. Using the client's style information collection Get two re-styled global auxiliary data 、 , and through the original global data and restyled global auxiliary data 、 Simulate the unknown target domain to score the generalization ability of the client local model, and then optimize the aggregation weight of the client local model based on the generalization ability score to strengthen the contribution of clients with high generalization ability and improve the model aggregation effect. 3) Model aggregation. Assume that global aggregation is performed Round, in In the iterative global training, the server Aggregate the client local model to obtain the global model of the t+1th iteration , and returns it to the client for the next round of iterative training. Finally, the server applies the global model obtained after multiple rounds of training to the unknown target domain.
[0074] The training steps of the image classification model specifically include:
[0075] Step 1: Before training begins, each client calculates its historical medical image style. The server collects the client's historical medical image style and broadcasts it to all clients. Afterwards, the client and server perform image style transfer on the original training dataset based on the global style information to generate a multi-style training dataset. Specifically:
[0076] The AdaIN model is used as a style transfer method. As a real-time arbitrary style transfer model, it can achieve style transfer efficiently and directly. In addition, existing research shows that the style information it shares cannot be used to reconstruct private data, which meets the privacy protection requirements of federated learning. Specifically, for the content images in the client's historical medical images, and style images , use VGG (Visual Geometry Group Network) encoder to obtain content features and style characteristics :
[0077] ;
[0078] in represents the VGG encoder, and ,in Represents the sum of the dimensions of feature batch, channel, height, and width. Further, the mean of each feature can be obtained and standard deviation , the calculation formula is:
[0079] ;
[0080] ;
[0081] in ; Represents the dimensions of feature batch, channel, height, and width.
[0082] During the style transfer process, AdaIN first de-stylizes the content image (normalizes the content features ), and then use the style image for stylization (using style features Perform affine transformation), and obtain the processed features ,Right now:
[0083] ;
[0084] Finally, the processed features Sent to the decoder Generate re-stylized images ,Right now:
[0085] .
[0086] Therefore, the client For example, its local historical medical image style The calculation process is as follows:
[0087] ;
[0088] in, represents the historical medical image style of the kth client; 、 denote the mean and standard deviation of the k-th client image feature respectively; Represents the mathematical expectation of the original private data; represents an image sample; represents the original private dataset of the kth client; is the result after data encoding; Indicates the calculation of the mean, Indicates the calculation of variance.
[0089] It can be seen from this that the client's local historical medical image style is the mathematical expectation of all local image styles, which can effectively represent the style of the client's global data.
[0090] like Figure 2 As shown, the server collects the local historical medical image styles of all clients and builds a global style set , and broadcast it to the client. The server fuses the local historical medical image styles to generate a fused global style, namely:
[0091] ;
[0092] in, Indicates the number of clients; Indicates the integration of global style; represents the style fusion weight of the k-th client, and . Then, as Figure 2As shown, the re-stylized global auxiliary data is generated according to the fusion global style, that is, the fusion global style is used as the global migration style, and the original global dataset is Perform two style transfers to obtain a global auxiliary dataset and , after integration, we can get the global data set .
[0093] Each client , which will be drawn from the global style set Randomly select the styles of the other two data domains as the privacy transfer styles, and then Figure 2 As shown, for the original privacy dataset Perform two style transfers to obtain a privacy-assisted dataset and ,in and Original privacy data The results after re-stylization. These two types of datasets will be combined into a re-stylized privacy dataset However, there is a special case where there are only two clients participating in the training, and the client can only select the style information of the other client for data enhancement. This example assumes that the number of clients participating in federated learning is .
[0094] Step 2: The server initializes the global model and sends it to the client. The client then trains a local model on its private dataset. Regularization is used during the training process to prevent the client's local model from overfitting to the local data domain. Specifically:
[0095] During local training, the JS divergence consistency loss is introduced as a regularization term in the objective function of the client-side local model. The core idea is to strengthen the consistency of the model's prediction results for the original image and its corresponding re-stylized image, thereby enabling the model to learn domain-invariant features. The objective function of the client-side local model is expressed as:
[0096] ;
[0097] in, represents the objective function of the local model of the kth client in the tth iteration; represents the k-th client privacy dataset; represents the k-th client local model in the t-th iteration; Represents the mathematical expectation on the private dataset; represents the cross entropy loss function; Indicates calculation of JS divergence; Indicates the strength of JS consistency restrictions; represents the predicted distribution of the original private data of the kth client in the tth iteration; 、 represents the predicted distribution of the k-th client's privacy auxiliary data in the t-th iteration.
[0098] Finally, the client local model is updated according to the following formula:
[0099] ;
[0100] in, Indicates the k-th client local model in the t+1th iteration; Indicates the update model; represents the k-th client local model in the t-th iteration; Indicates the client local learning rate; Indicates the gradient is obtained; represents the local objective function of the k-th client.
[0101] Step 3: After local training is completed, the client uploads the client's local model to the server. The server evaluates the generalization ability of each client's local model on the global dataset, optimizes the aggregation weight based on the generalization ability score of the client's local model, and finally aggregates the local models to obtain the image classification model, which is specifically:
[0102] The generalization ability of each client-side local model can be quantified by the consistency of its predictions on the original global data and the re-stylized global auxiliary data, which can be expressed as follows:
[0103] ;
[0104] in, represents the generalization ability score of the k-th client local model in the t-th iteration; represents the predicted distribution of the original global data in the tth iteration; 、 represents the predicted distribution of the global auxiliary data in the tth iteration; represents the KL divergence.
[0105] The generalization ability of the local model is positively correlated with the consistency of its predictions on the original global data and the re-stylized global auxiliary data. The higher the generalization ability score, the stronger the model's generalization ability.
[0106] When optimizing the client-side local model aggregation weight, the client data volume ratio and the model generalization capability are comprehensively considered to determine its aggregation weight, namely:
[0107] ;
[0108] right Normalize and you can get the final aggregation weight:
[0109] ;
[0110] in, represents the initial aggregate weight of the k-th client in the t-th iteration; represents the influencing factor of generalization ability score; represents the final aggregate weight of the k-th client in the t-th iteration.
[0111] Finally, the client local model is weightedly aggregated according to the aggregation weight to obtain a trained image classification model. That is, the image classification model update formula is expressed as:
[0112] ;
[0113] in, Represents the t+1th round of iterative image classification model.
[0114] This implementation utilizes image style transfer technology to enrich the domain diversity of the client's local data and introduces a prediction consistency regularization term into the optimization objective. This allows the model to maintain stable output when presented with both the original sample and its stylized version, alleviating the model's overfitting to the local data domain while helping it learn domain-invariant features. Furthermore, the generalization performance of the client's local model is tested by simulating an unknown domain environment using multi-style versions of global data. Based on the generalization performance evaluation results, the contribution ratio of the client's local model to the image classification model is dynamically adjusted, further improving aggregation and model generalization performance.
[0115] Example 2
[0116] Based on Example 1, this example introduces an experimental example of a medical image classification method generalized for federated domains:
[0117] This example uses the PACS and Office-Home datasets. PACS has four different image style domains (Photo, Art, Cartoon, and Sketch), while Office-Home also contains images from four different domains (Art, Clipart, Product, and Real-World). During training, if one dataset is used as the client training dataset, a portion of data from the other dataset is selected as the global dataset to assist in training. Figure 3-6 Shows the results of style migration of data. Figure 3 、 Figure 4They represent the results of PACS and Office-Home re-stylization, respectively. Each row represents the domain to which the image content belongs, and each column represents the domain to which the image style belongs. The images on the main diagonal are the original data. Figure 5 、 Figure 6 They represent the re-stylized results when PACS and Office-Home are used as global datasets, respectively. The first column is the original data, and the last two columns are the re-stylized images.
[0118] The experiment uses the pre-trained AdaIN model for image style transfer, and uses the ImageNet pre-trained ResNet50 and ResNet18 network models for training on the PACS and Office-Home datasets, respectively. It is carried out in the form of LODO (Leave One Domain Out), that is, a specific domain in the dataset is selected as the unknown target test domain, and the data of the remaining domains are deployed to each client as distributed training source domains. For the PACS and Office-Home datasets, the training batch size for each client is 64 and 32, respectively, and a stochastic gradient descent optimizer with a learning rate of 1e-3, a momentum of 0.9, and a weight decay of 1e-4 is used. Total communication rounds Set to 40, the number of local training iterations The size of the public dataset, the strength of the JS consistency restriction, and the degree of influence of the model generalization ability on the aggregation weight are set to 3000, 9, and 0.8, respectively (i.e. , , ).
[0119] The experiment selected the traditional federated learning algorithm FedAvg as a benchmark solution and introduced multiple methods for comparison. Specifically, it included three domain generalization methods for centralized scenarios, namely JiGen, RSC, and MixStyle, and two federated domain generalization methods, FedDG and CCST. Among them, JiGen and RSC can be directly integrated into FedAvg without additional modification. When adapting to federated learning scenarios, MixStyle needs to be adapted to the client-side internal version. That is, the data within the client is treated as an independent "domain". By performing style fusion operations on the internal images of the client, it can be effectively integrated with the federated learning framework.
[0120] Table 1 presents the test results of our method and comparison schemes on the PACS and Office-Home datasets. Each column represents the test accuracy when the domain is used as an unknown target domain, and the final result is the average of the model's test results in the last 10 rounds. As can be seen from the table, the proposed method outperforms other schemes on both datasets, with an average test accuracy that is more than 1.8% higher than that of other schemes. Specifically, on the PACS dataset, the proposed scheme achieved an average accuracy of 85.01%, 2.66% higher than the next best scheme. Similarly, on the Office-Home dataset, our scheme performed best, with an average accuracy that was 1% higher than the next best scheme.
[0121] Table 1 Test results of different methods on PACS and Office-Home datasets
[0122]
[0123] Overall, this method outperforms other approaches in enhancing model generalization. This is because it leverages image style transfer technology to enrich the diversity of the client-side training data domain. Diversified training data helps the client-side local model learn more generalizable feature representations, effectively mitigating overfitting to the local raw data and ultimately improving model generalization. However, centralized domain generalization solutions like JiGen and RSC only have access to a single source domain data in federated learning, making them prone to overfitting. Consequently, local models trained on different source domains, when aggregated, struggle to adapt well to the unknown target domain. MixStyle, FedDG, and CCST share similarities with this approach, all aiming to improve model generalization by diversifying the feature distribution of training data. However, within the federated learning framework, MixStyle only assumes that images within the data domain are independent "domains" for style fusion, which has limited effectiveness in improving data diversity. FedDG utilizes amplitude information in the image frequency space as style information and implements amplitude exchange to diversify data distribution between clients. However, the degree of image change caused by amplitude exchange is relatively small, which is far from enough for data with large differences between domains. CCST achieves consistency in client data distribution by transferring the style of the remaining client images to the local area, which to a certain extent expands the diversity of the data domain. However, in reality, this approach makes the client data domain converge, which will cause the client local model to overfit the data domain, which is not conducive to the generalization of the global model to unknown target domains. In contrast, this method uses the style-transferred data for local regularization. By strengthening the consistency of the model's prediction results for the original image and its re-stylized image, it encourages the model to learn domain-invariant features and enhances its own generalization ability.
[0124] Figure 7-10 and Figure 11-14 The test accuracy curves of different schemes on PACS and Office-Home datasets are presented. Each figure shows the results when a specific domain is used as the test data domain. As can be seen from the figure, the convergence speed of this method is significantly faster than that of other comparison schemes. Figure 9 For example, this method converged within 5 rounds. This advantage is mainly attributed to two aspects: first, the deployment of local regularization enables the model to avoid overfitting and thus prevent the loss of generalization performance; second, the use of the aggregation weight optimization mechanism allows the local model with strong generalization capabilities to play a greater role in the aggregation process, further accelerating the convergence process.
[0125] To demonstrate the detailed contributions of each component in this approach, ablation experiments were conducted on two datasets. Specifically, this approach can be divided into two parts: Local Regularization (LR) and Aggregation Weight Optimization (AWO). The experiments used these two components to test their impact on the final model performance, with the results shown in Table 2. As can be seen, deploying both LR and AWO significantly improves model performance. Their collaboration ultimately achieves a test accuracy of 75.06% on both test datasets, surpassing the baseline FedAvg by 12.19% and exceeding the performance of the LR and AWO components deployed separately by 0.26% and 3.4%, respectively. Furthermore, the experimental results demonstrate that LR significantly outperforms AWO in improving model performance. This is because LR directly enhances the generalization capability of each client's local model by helping the model learn domain-invariant feature representations, effectively alleviating overfitting. AWO, as a server-side optimization strategy, relies on the quality of the client's local model, making LR a more critical factor in improving overall system performance.
[0126] Table 2 Test results based on different components on different datasets
[0127]
[0128] Example 3
[0129] Based on Example 1 or 2, this embodiment introduces a medical image classification system for federated domain generalization, including:
[0130] An image acquisition module is used to: acquire medical images to be classified;
[0131] An image classification module is used to: input the medical image to be classified into a pre-trained image classification model and output a medical image classification result;
[0132] A model training module is used for: wherein the training of the image classification model includes:
[0133] Obtain historical medical image styles for each client;
[0134] Generate a global style set based on the historical medical image styles of each client, and use the global style set to perform image style transfer on the original training dataset to generate a multi-style training dataset; the multi-style training dataset includes a private dataset and a global dataset;
[0135] The image classification model to be trained is distributed to each client as the client-side local model to be trained. The private dataset is input into the client-side local model to obtain a trained client-side local model. JS divergence consistency loss is added to the objective function of the client-side local model as a regularization term.
[0136] Receive the trained local models of each client, calculate the generalization ability score of each client local model based on the global data set, calculate the aggregation weight of the client local model based on the generalization ability score, and perform weighted aggregation on the client local models according to the aggregation weight to obtain a trained image classification model.
[0137] The specific functional implementation of each of the above modules can be found in the relevant content of the method in Example 1 and will not be elaborated on here.
[0138] Example 4
[0139] This embodiment introduces a computer-readable storage medium on which a computer program is stored. When the program is executed by a processor, the steps of the medical image classification method for federated domain generalization as described in Example 1 or 2 are implemented.
[0140] Those skilled in the art will appreciate that the embodiments of the present application may be provided as methods, systems, or computer program products. Therefore, the present application may take the form of an entirely hardware embodiment, an entirely software embodiment, or an embodiment combining software and hardware. Furthermore, the present application may take the form of a computer program product implemented on one or more computer-usable storage media (including but not limited to magnetic disk storage, CD-ROM, optical storage, etc.) containing computer-usable program code.
[0141] The present application is described with reference to the flowcharts and / or block diagrams of the methods, devices (systems), and computer program products according to the embodiments of the present application. It should be understood that each process and / or block in the flowchart and / or block diagram, as well as the combination of processes and / or blocks in the flowchart and / or block diagram, can be implemented by computer program instructions. These computer program instructions can be provided to a processor of a general-purpose computer, a special-purpose computer, an embedded processor, or other programmable data processing device to produce a machine, so that the instructions executed by the processor of the computer or other programmable data processing device generate instructions for implementing the processes in the flowchart and / or block diagram. Figure 1 a process or multiple processes and / or boxes Figure 1 A device that provides the functions specified in a block or multiple blocks.
[0142] These computer program instructions may also be stored in a computer readable memory that can direct a computer or other programmable data processing device to work in a specific manner, so that the instructions stored in the computer readable memory produce an article of manufacture comprising an instruction device, which implements the process Figure 1 a process or multiple processes and / or boxes Figure 1 The function specified in one or more boxes.
[0143] These computer program instructions can also be loaded onto a computer or other programmable data processing device so that a series of operational steps are executed on the computer or other programmable device to produce a computer-implemented process, thereby providing the instructions executed on the computer or other programmable device for implementing the process. Figure 1 a process or multiple processes and / or boxes Figure 1 A step that specifies a function in one or more boxes.
[0144] The embodiments of the present invention are described above in conjunction with the accompanying drawings, but the present invention is not limited to the above-mentioned specific implementation methods. The above-mentioned specific implementation methods are merely illustrative and not restrictive. Under the guidance of the present invention, ordinary technicians in this field can also make many forms without departing from the scope of protection of the purpose of the present invention and the claims, which are all protected by the present invention.
Claims
1. A medical image classification method for federated domain generalization, characterized by: include: Obtaining medical images to be classified; Inputting the medical image to be classified into a pre-trained image classification model and outputting a medical image classification result; The training of the image classification model includes: Obtain historical medical image styles for each client; Generate a global style set based on the historical medical image styles of each client, and use the global style set to perform image style transfer on the original training dataset to generate a multi-style training dataset; the multi-style training dataset includes a private dataset and a global dataset; The image classification model to be trained is distributed to each client as the client-side local model to be trained. The private dataset is input into the client-side local model to obtain a trained client-side local model. JS divergence consistency loss is added to the objective function of the client-side local model as a regularization term. Receive the trained local models of each client, calculate the generalization ability score of each client local model based on the global data set, calculate the aggregation weight of the client local model based on the generalization ability score, and perform weighted aggregation on the client local models according to the aggregation weight to obtain a trained image classification model.
2. The federated generalization-oriented medical image classification method according to claim 1, characterized in that: Use the global style set to perform image style transfer on the original training dataset to generate a multi-style training dataset, including: Obtaining an original training dataset; the original training dataset includes an original private dataset and an original global dataset; Randomly select the historical medical image style of other clients in the global style set as the privacy transfer style, perform style transfer on the original private dataset according to the privacy transfer style, and obtain a privacy auxiliary dataset. Then, obtain the private dataset based on the local private dataset and the privacy auxiliary dataset. The historical medical image styles of the client in the global style set are fused to obtain a fused global style; the fused global style is used as the global migration style, and the original global dataset is style-migrated according to the global migration style to obtain a global auxiliary dataset; and a global dataset is obtained based on the original global dataset and the global auxiliary dataset.
3. The federated generalization-oriented medical image classification method according to claim 2, characterized in that: The historical medical image styles of the client in the global style set are integrated to obtain a fused global style, including: ; ; ; in, represents the historical medical image style of the kth client; 、 denote the mean and standard deviation of the k-th client image feature respectively; Represents the mathematical expectation of the original private data; represents an image sample; represents the original private dataset of the kth client; is the result after data encoding; Indicates the calculation of the mean, Indicates the calculation of variance; Represents a global style set; Indicates the number of clients; Indicates the integration of global style; represents the style fusion weight of the k-th client.
4. The medical image classification method for federated domain generalization according to claim 1, characterized in that: The objective function of the client local model is expressed as: ; in, represents the objective function of the local model of the kth client in the tth iteration; represents the k-th client privacy dataset; represents the k-th client local model in the t-th iteration; Represents the mathematical expectation on the private dataset; represents the cross entropy loss function; Indicates calculation of JS divergence; Indicates the strength of JS consistency restrictions; represents the predicted distribution of the original private data of the kth client in the tth iteration; 、 represents the predicted distribution of the k-th client's privacy auxiliary data in the t-th iteration; represents the original private data of the kth client; 、 represents the k-th client privacy auxiliary data; Indicates the label corresponding to the original private data of the kth client.
5. The medical image classification method for federated domain generalization according to claim 1, characterized in that: The updating process of the client local model is expressed as: ; in, Indicates the k-th client local model in the t+1th iteration; Indicates the update model; represents the k-th client local model in the t-th iteration; Indicates the client local learning rate; Indicates the gradient is obtained; represents the local objective function of the kth client; represents the k-th client private dataset.
6. The medical image classification method for federated domain generalization according to claim 1, characterized in that: The generalization ability score of each client local model is expressed as: ; in, represents the generalization ability score of the k-th client local model in the t-th iteration; represents the predicted distribution of the original global data in the tth iteration; 、 represents the predicted distribution of the global auxiliary data in the tth iteration; Represents original global data; 、 Represents global auxiliary data; represents the k-th client local model in the t-th iteration; Indicates the number of original global data; represents the KL divergence.
7. The medical image classification method for federated domain generalization according to claim 1, characterized in that: Calculate the aggregation weight of the client local model based on the generalization ability score, including: ; ; in, represents the initial aggregate weight of the k-th client in the t-th iteration; represents the amount of local private data of the kth client; represents the influencing factor of generalization ability score; represents the generalization ability score of the k-th client local model in the t-th iteration; represents the final aggregate weight of the k-th client in the t-th iteration; Indicates the number of clients.
8. The medical image classification method for federated domain generalization according to claim 1, characterized in that: The trained image classification model is expressed as: ; in, Represents the t+1th round of iterative image classification model; represents the final aggregate weight of the k-th client in the t-th iteration; represents the k-th client local model in the t-th iteration; Indicates the number of clients.
9. A medical image classification system for federated domain generalization, characterized in that: include: An image acquisition module is used to: acquire medical images to be classified; An image classification module is used to: input the medical image to be classified into a pre-trained image classification model and output a medical image classification result; A model training module is used for: wherein the training of the image classification model includes: Obtain historical medical image styles for each client; Generate a global style set based on the historical medical image styles of each client, and use the global style set to perform image style transfer on the original training dataset to generate a multi-style training dataset; the multi-style training dataset includes a private dataset and a global dataset; The image classification model to be trained is distributed to each client as the client-side local model to be trained. The private dataset is input into the client-side local model to obtain a trained client-side local model. JS divergence consistency loss is added to the objective function of the client-side local model as a regularization term. Receive the trained local models of each client, calculate the generalization ability score of each client local model based on the global data set, calculate the aggregation weight of the client local model based on the generalization ability score, and perform weighted aggregation on the client local models according to the aggregation weight to obtain a trained image classification model.
10. A computer-readable storage medium having a computer program stored thereon, characterized in that: When the program is executed by a processor, the steps of the medical image classification method for federated domain generalization as described in any one of claims 1 to 8 are implemented.