Method for Sharing Medical Institution Data Based on Federated Learning
By building structural knowledge extraction module, site collaboration enhancement module and local update module, the federated learning model performance degradation caused by data differences between different medical institutions is solved, and performance improvement and training speed are improved.
Patent Information
- Application Number
- CN202311105204.7
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2023-08-30
- Publication Date
- 2025-05-30
- Estimated Expiration
- 2043-08-30
AI Technical Summary
Data differences between different medical institutions lead to a degradation in the performance of models established by federated learning methods, affecting the performance of the federal global model.
By building a federal medical data structure knowledge extraction module, a site collaboration enhancement module and a local update module, we eliminate data characteristics differences, dynamically weight the structural knowledge, obtain consensus structure knowledge, and update model parameters locally to improve model performance.
This improves the performance of the federated learning model, improves the training speed, and solves the problem of model performance degradation caused by data differences.
Smart Images

Figure CN117171267B_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the technical field of federated learning, and particularly relates to a method for sharing medical institution data based on federated learning. Background Art
[0002] With the development of deep learning, big data technology can use data widely distributed at various sites (such as major medical institutions) to establish an efficient model to improve the feasibility of artificial technology in related industries. However, laws and regulations on personal information protection restrict the use of decentralized data widely distributed. Therefore, the feasibility of directly collecting medical data for centralized model establishment is restricted.
[0003] The general framework of federated learning proposed by McMahan et al. uses decentralized and widely distributed data to train a federated global model; MJ Sheller et al. introduced federated learning into the use of medical data, providing a feasible solution for decentralized collaborative modeling between medical institutions; different from the traditional centralized training method that directly uses data from multiple sites, the federated learning method establishes a federated model through model aggregation, only requiring the exchange of model parameters between the server and the sites; the main purpose of federated learning is to establish a stable federated global model. The main steps for establishing a joint model include: first, the sites log in to the server and send the global model parameters to the sites. Then, each site participating in the training uses its own local private data to train to obtain new model parameters and sends them back to the server. Finally, the server aggregates the updated local models sent back by all sites to finally obtain a new federated global model. The above steps are the general framework of federated learning.
[0004] Applying the above general framework of federated learning to actual medical data requires considering expanding the framework to adapt to the data set; specifically, the data differences between different medical institutions may lead to a problem of performance degradation of the model established by the federated learning method, and the factors causing data differences include different data collection protocols, data sampling devices, and data processing methods used by different medical institutions; the data differences result in differences in data feature distribution, image appearance, etc. between site data, thereby affecting the performance of the federated global model established by federated learning. Summary of the Invention
[0005] The purpose of the present invention is to provide a method for sharing medical institution data based on federated learning with improved performance and training speed.
[0006] The method for sharing medical institution data based on federated learning provided by the present invention includes the following steps:
[0007] S1. Obtain a public medical dataset and establish the association between medical sites and service terminals participating in federated training;
[0008] S2. Using the public medical dataset obtained in step S1, construct a federated medical data structure knowledge extraction block to extract a set of structural knowledge;
[0009] S3. Using the set of structural knowledge obtained by extraction in step S2, construct a site collaboration enhancement module for dynamic weighting processing to obtain consensus structural knowledge;
[0010] S4. Using the public medical dataset obtained in step S1 and the set of structural knowledge obtained in step S3, construct a local update module to calculate the local overall loss;
[0011] S5. Using the local overall loss calculated in step S4, update the local model to complete the training process of the site;
[0012] S6. For the sites that have completed training in step S5, perform aggregation processing on the server side to obtain a new global model of federated medical data;
[0013] S7. Repeat the above steps S2 - S6, perform multiple rounds of federated training in a loop until the global model of federated medical data obtained reaches the desired performance, and complete the construction of the federated medical data model;
[0014] S8. Using the federated medical data model constructed in step S7, realize data sharing among medical institutions;
[0015] The obtaining of the public medical dataset described in step S1 and the establishment of the association between medical sites and service terminals participating in federated training specifically include:
[0016] The public medical dataset includes a labeled dataset, an unlabeled dataset, and a third - party public dataset with differences in features and data volume from all sites participating in federated training; define the public medical dataset as D 0 ;
[0017] Establish a global model θ on the server side and initialize the training environment, including the initial parameters of the model and the initial values of the attention map set; define that there are K sites in the federated learning system, each site is represented as k, and the initialized local model of the site is Initialize the local model of the previous federated training as
[0018] The server samples a new public sample X 0 ~D 0 from the public medical dataset, where X 0 represents the public sample;
[0019] Using the public medical dataset obtained in step S1 described in step S2, construct a federal medical data structure knowledge extraction block to extract a set of structural knowledge, specifically including:
[0020] (2-1) The site-local model θ k The predicted value output for the feature map The value of the average gradient of on the channel c is defined as β c , and calculate β using the following formula c :
[0021]
[0022] Among them, H 1 represents the total longitudinal pixel length of the feature map M; W 1 represents the total transverse pixel length of the feature map M; i represents the i-th transverse pixel of the feature map M; j represents the j-th longitudinal pixel of the feature map M; y c represents the predicted value output by the local model θ k ; M c represents the value of the feature map on the channel c;
[0023] In a deep learning neural network driven by medical image data, select the bottommost module of the encoder part in U-Net as the feature map M;
[0024] (2-2) Aggregate the average gradient calculated in step (2-1) with the feature map, filter the unresponsive part, and obtain the structural knowledge. The calculation formula is as follows:
[0025]
[0026] Among them, A c represents the value of the structural knowledge on the channel c; C represents the total number of channels; c represents the channel with the largest amount of labeled data; ReLU(·) represents a non-linear activation function, and the calculation formula is as follows:
[0027]
[0028] Use the obtained structural knowledge to replace the parameters of the model as the medium for knowledge transfer between sites, reducing the inconsistency of semantic knowledge captured by models at different sites;
[0029] For site k ∈ K, define the corresponding output as
[0030] Use the following defined function to represent the structural knowledge module in steps (2-1) and (2-2):
[0031] φ(X,f(X;θ))
[0032] Among them, X represents the sample input to the function; f(·) represents the model used in the inference process; θ represents the parameters corresponding to the model;
[0033] Using the common sample X 0 , the local models of all K sites On the server side, through the constructed knowledge structure module φ(X, f(X; θ)), the structural knowledge set is extracted Among them, φ(X 0 ; f(X 0 ; θ k )) represents the structural knowledge module function extracted based on the sample X 0 and the model θ of site k k ;
[0034] Using the structural knowledge set extracted in step S2 described in step S3, construct a site collaboration enhancement module for dynamic weighting processing to obtain consensus structural knowledge, specifically including:
[0035] (3-1) Define the model θ of site k k The inference output on the common sample X 0 is Z k = f(X 0 ; θ k ); Collect the inference outputs of all K sites, defined as
[0036] (3-2) Use the following formula to represent the updated weighted vector W of site k k :
[0037]
[0038] Among them, W′ k represents the updated weighted vector of the site; η 1 represents the learning rate of the loss function; W i represents the weighted weight vector set for each site, with a length of K; ▽ represents the Nabla operator, used to calculate the derivative; represents the mean square error loss function; W kj represents the weight value of site j for the target site k; Z j represents the predicted value of the site j model output on the common sample; ρ represents the control parameter, used to limit the weight;
[0039] (3-3) Repeat the updated weighted vector W′ obtained in step (3-2) k , a total of K times, and update the weight W′ corresponding to site k each time k , and finally obtain the updated weight set The set of structural knowledge obtained by extraction in step S2 is aggregated to obtain the consensus structural knowledge A for site k t , and A is calculated using the following formula t :
[0040]
[0041] where represents the structural knowledge extracted by site j, and the calculation formula is as follows:
[0042]
[0043] Step S4 constructs a local update module using the public medical dataset obtained in step S1 and the set of structural knowledge obtained in step S3, and calculates the local overall loss, which specifically includes:
[0044] (4-1) For each medical site k participating in training, download the global model θ, the consensus structural knowledge A t , and the new public sample X 0 to the local site, and the site aligns the feature space of the site model using the sample and the consensus structural knowledge;
[0045] (4-2) Site k updates the parameters of the local model:
[0046] The following formula is used to represent the local loss function of the site
[0047]
[0048] where α represents the proportionality coefficient; represents the task loss function; represents the local integrated knowledge distillation loss function;
[0049] (4-3) Using the local dataset D k ={(X i ,y i )}, define the task loss function as follows:
[0050]
[0051] where X i represents the i-th batch of samples in the local data; y i represents the label corresponding to the i-th batch of samples in the local data; f(X i ;θ k ) represents the output of the local model; loss represents the function, defined as: In the medical segmentation task, Represents the sample space of data, and the corresponding label space is where C represents the number of channels of the sample image; H represents the height of the sample image; W represents the width of the sample image;
[0052] The task loss function is calculated using the following formula:
[0053]
[0054] where y represents the label corresponding to the sample; Represents the predicted result, and the calculation formula is as follows:
[0055]
[0056] (4-4) The local structure knowledge integration loss includes: Represents the similarity between the attention map extracted by the local model and the average value of the attention map set downloaded from the server, Represents the similarity between the attention map extracted by the local model and the attention map extracted by the local model in the previous communication round;
[0057] The local structure knowledge integration loss is calculated using the following formula
[0058]
[0059] where, Represents the structural knowledge output by the current local model θ to be updated k The calculation formula is as follows:
[0060]
[0061] where, φ(X 0 ; f(X 0 ; θ i )) represents the structural knowledge module function extracted based on the sample X 0 and the model θ i ;
[0062] A t Represents the consensus structural knowledge obtained by downloading from the server;
[0063] Represents the structural knowledge output by the local model obtained based on the previous round of federated training, and the calculation formula is as follows:
[0064]
[0065] where, Indicates based on sample X 0 and the model The structural knowledge module function extracted;
[0066] τ represents the temperature coefficient, which is used to control the scope covered by the structural knowledge;
[0067] 0≤sim(p,q)≤1 represents the measurement function of the similarity between the inputs p and q. The closer the value of the function is to 1, the more similar p and q are;
[0068] For medical data, the Dice coefficient is used to implement the similarity function, and the calculation formula is as follows:
[0069]
[0070] Using the local overall loss calculated in step S4 in step S5 to update the local model and complete the training process of the site, specifically including:
[0071] For any site k, the update process is represented by the following formula:
[0072]
[0073] Among them, θ k represents the old local model of site k; θ′ k represents the updated local model of site k; η 2 represents the learning rate of the task loss of; η 3 represents the learning rate of the local integrated knowledge distillation loss function of; α represents the proportional coefficient of the two losses; represents the gradient used to update the local model;
[0074] For any site k, store the currently updated federated model and use it for the next parameter update, which is represented by the following formula:
[0075]
[0076] Among them, after site k completes training, the updated local model θ′ k , is assigned to as the local model stored by site k in the previous communication round; then, the client submits the updated local model θ′ k to the server; the server waits for all K sites to submit their local models to obtain the model set
[0077] The sites that have completed training in step S6 as described in step S6 are aggregated on the server to obtain a new global model of federated medical data, specifically including:
[0078] Aggregate the local model parameter sets submitted by all K sites on the server side Obtain a new global model of federated medical data Among them, p k represents the weighted coefficient of the site, and the calculation formula is as follows:
[0079]
[0080] This method for sharing medical institution data based on federated learning provided by the present invention eliminates the problem of performance degradation of the federated global model caused by the data feature difference problem faced in the cooperation of multiple medical sites to establish a joint model by constructing a federated medical data structure knowledge extraction module, a site collaboration enhancement module, and a local update module; the performance of the present invention is improved and the training speed is increased. Brief Description of the Drawings
[0081] Figure 1 It is a schematic flowchart of the method of the present invention.
[0082] Figure 2 It is a schematic diagram of the site collaboration enhancement module of the method of the present invention.
[0083] Figure 3 It is a schematic diagram of the action of the structure knowledge module of the present invention on the site model update.
[0084] Figure 4 It is a schematic diagram of the visualization demonstration result in the embodiment of the method of the present invention. Detailed Embodiment
[0085] As Figure 1 shown is a schematic flowchart of the method of the present invention: This method for sharing medical institution data based on federated learning provided by the present invention includes the following steps:
[0086] S1. Obtain a public medical data set and establish an association between the medical sites participating in the federated training and the service terminal; specifically including:
[0087] The public medical data set includes a labeled data set, an unlabeled data set, and a third-party public data set with differences in features and data volume from all the sites participating in the federated training; define the public medical data set as D 0 ;
[0088] Establish a global model θ on the server side and initialize the training environment, including the initial parameters of the model and the initial values of the attention map set; define that there are K sites in the federated learning system, each site is represented as k, and the initialized local model of the site is Initialize the local model of the previous federated training as
[0089] The server samples a new public sample X from the public medical dataset 0 ~D 0 ; where X 0 represents the public sample;
[0090] S2. Using the public medical dataset obtained in step S1, construct a federated medical data structure knowledge extraction block to extract a set of structural knowledge; specifically including:
[0091] (2-1) The average gradient of the predicted value output by the local model θ k of the feature map on channel c is defined as β c , and β is calculated using the following formula c :
[0092]
[0093] where H 1 represents the total number of vertical pixels of the feature map M; W 1 represents the total number of horizontal pixels of the feature map M; i represents the i-th horizontal pixel of the feature map M; j represents the j-th vertical pixel of the feature map M; y c represents the predicted value output by the local model θ k ; M c represents the value of the feature map on channel c;
[0094] In a deep learning neural network driven by medical image data, select the bottommost module of the encoder part in U-Net as the feature map M;
[0095] (2-2) Aggregate the average gradient calculated in step (2-1) with the feature map, filter the unresponsive part, and obtain the structural knowledge. The calculation formula is as follows:
[0096]
[0097] where A c represents the value of the structural knowledge on channel c; C represents the total number of channels; c represents the channel with the largest amount of labeled data; ReLU(·) represents a non-linear activation function, and the calculation formula is as follows:
[0098]
[0099] Use the obtained structural knowledge to replace the parameters of the model as the medium for knowledge transfer between sites, reducing the inconsistency of semantic knowledge captured by models at different sites;
[0100] For site k ∈ K, define the corresponding output as
[0101] Use the following defined function to represent the structural knowledge modules in steps (2-1) and (2-2):
[0102] φ(X, f(X; θ))
[0103] where X represents the sample input to the function; f(·) represents the model used in the inference process; θ represents the parameters corresponding to the model;
[0104] Use the common sample X 0 , the local models of all K sites At the server side, through the constructed structural knowledge module φ(X, f(X; θ)), extract the structural knowledge set where φ(X 0 ; f(X 0 ; θ k )) represents the structural knowledge module function extracted based on the sample X 0 and the model θ of site k k ;
[0105] S3. Use the structural knowledge set obtained by extraction in step S2 to construct a site collaboration enhancement module for dynamic weighting processing to obtain consensus structural knowledge; specifically including:
[0106] As Figure 2 shown is a schematic diagram of the site collaboration enhancement module of the method of the present invention:
[0107] (3-1) Define the model θ of site k k The inference output on the common sample X 0 is Z k = f(X 0 ; θ k ); Collect the inference outputs of all K sites and define them as
[0108] (3-2) Use the following formula to represent the update of the weighted vector W of site k k :
[0109]
[0110] where W′ k represents the weighted vector of the updated site; η 1 represents the learning rate of the loss function; W i represents the weighted weight vector set for each site, with a length of K; represents the Nabla operator, used to calculate the derivative; represents the mean square error loss function; W kjDenote the weight value of site j for the target site k; Z j Denote the predicted value output by the model of site j on the public samples; ρ represents the control parameter used to limit the weight;
[0111] (3-3) The weighted vector W′ of the updated site obtained by repeating step (3-2) k , repeat a total of K times, and update the weight W′ corresponding to site k each time k , and finally obtain the updated weight set For the set of structural knowledge extracted in step S2 Perform an aggregation process to obtain the consensus structural knowledge A for site k t , calculate A using the following formula t :
[0112]
[0113] where Denote the structural knowledge extracted by site j, and the calculation formula is as follows:
[0114]
[0115] S4. Use the public medical dataset obtained in step S1 and the set of structural knowledge obtained in step S3 to construct a local update module and calculate the local overall loss; specifically including:
[0116] (4-1) For each medical site k participating in the training, download the global model θ, the consensus structural knowledge A t , and the new public samples X 0 to the local site, and the site uses the samples and the consensus structural knowledge to align the feature space of the site model;
[0117] (4-2) Site k updates the parameters of the local model:
[0118] The local loss function of the site is represented by the following formula
[0119]
[0120] where α represents the proportionality coefficient; Denote the task loss function; Denote the local integrated knowledge distillation loss function;
[0121] (4-3) Use the local dataset D k ={(X i ,y i )}, define the task loss function as follows:
[0122]
[0123] Among them, X i represents the i-th batch of samples in the local data; y i represents the label corresponding to the i-th batch of samples in the local data; f(X i ; θ k ) represents the output of the local model; loss represents a function, defined as: In the medical segmentation task, represents the sample space of the data, and the corresponding label space is Among them, C represents the number of channels of the sample image; H represents the height of the sample image; W represents the width of the sample image;
[0124] The task loss function is calculated using the following formula:
[0125]
[0126] Among them, y represents the label corresponding to the sample; represents the predicted result, and the calculation formula is as follows:
[0127]
[0128] (4-4) The local structural knowledge integration loss includes: represents the similarity between the attention map extracted by the local model and the average value of the attention map set downloaded from the server, represents the similarity between the attention map extracted by the local model and the attention map extracted by the local model in the previous communication round;
[0129] The local structural knowledge integration loss is calculated using the following formula
[0130]
[0131] Among them, represents the structural knowledge output by the current local model θ to be updated k , and the calculation formula is as follows:
[0132]
[0133] Among them, φ(X 0 ; f(X 0 ; θ i )) represents the structural knowledge module function extracted based on the sample X 0 and the model θ i ;
[0134] At It represents the consensus structure knowledge obtained by downloading through the server.
[0135] It represents the local model obtained based on the previous federated training The output structure knowledge, and the calculation formula is as follows:
[0136]
[0137] Among them, It represents the structure knowledge module function extracted based on the sample X 0 and the model ;
[0138] τ represents the dimension coefficient, which is used to control the scope covered by the structure knowledge;
[0139] 0≤sim(p,q)≤1 represents the similarity measurement function of the inputs p and q. The closer the value of the function is to 1, the more similar p and q are;
[0140] For medical data, the Dice coefficient is used to implement the similarity function, and the calculation formula is as follows:
[0141]
[0142] S5. Using the local overall loss calculated in step S4, update the local model to complete the training process of the site; specifically including:
[0143] Such as Figure 3 shown is the schematic diagram of the action of the structure knowledge module of the method of the present invention on the site model update:
[0144] For any site k, the update process is represented by the following formula:
[0145]
[0146] θ k represents the old local model of site k; θ′ k represents the updated local model of site k; η 2 represents the learning rate of the task loss ; η 3 represents the learning rate of the local integrated knowledge distillation loss function ; α represents the proportional coefficient of the two losses; represents the gradient used to update the local model;
[0147] For any site k, store the currently updated federated model and use it for the next parameter update, which is represented by the following formula:
[0148]
[0149] Among them, after site k completes training, the updated local model θ′ k , is assigned as the local model stored by site k in the previous communication round; then, the client will submit the updated local model θ′ k to the server; the server waits for all K sites to submit their local models to obtain a model set
[0150] S6. For the sites that have completed training in step S5, perform aggregation processing on the server to obtain a new global model of federated medical data; specifically including:
[0151] Aggregate the local model parameter sets submitted by all K sites on the server Obtain a new global model of federated medical data where p k represents the weighted coefficient of the site, and the calculation formula is as follows:
[0152]
[0153] S7. Repeat the above steps S2 - S6, and perform multiple rounds of federated training in a loop until the global model of federated medical data obtained reaches the desired performance, and complete the construction of the federated medical data model;
[0154] S8. Use the federated medical data model constructed in step S7 to achieve data sharing among medical institutions.
[0155] The method of the present invention has been applied to multiple medical image segmentation tasks, including prostate image segmentation tasks, pathological tissue image segmentation tasks, multi-modal brain glioma segmentation tasks, and multi-site multi-device cardiac MRI image segmentation tasks; a joint medical image segmentation model is established through federated learning while protecting the privacy of data users; at the same time, the federated medical data structure knowledge extraction module, the site collaboration enhancement module, and the medical site local update module cooperate with each other. On the one hand, by exchanging structural knowledge between the server and the sites, knowledge can be migrated; on the other hand, taking the application of the present invention to the multi-modal brain glioma segmentation task as an example, the invention is exemplarily described. In this task, K = 8 medical sites are selected to participate in the federated learning training to establish a joint medical image segmentation model; since the data collection devices used by each medical site are different, the image features presented by different sites are different; the specific implementation process is described in detail in the following steps;
[0156] (1) Each site corresponds to a medical site i, and divides the private dataset D i into a training set and the test set The division ratio is 8:2;
[0157] The global model structure jointly established is a 2D U-Net network, and the parameters of the model are transmitted between the server and the sites;
[0158] Select the unlabeled dataset shared by other medical sites as the public dataset D 0 , and obtain the public sample X through sampling 0 ~D 0 ;
[0159] (2) The site collaboration enhancement module is applied to generate the consensus structure knowledge after dynamic weighting. In this embodiment, for the local model θ sent by each site k to the server k , the inference output of its public sample X 0 is Z k = f(X 0 ; θ k ), collect the inference outputs of all K sites, and define them as
[0160] (3) Use the following formula to represent the updated weighted vector W of site k k :
[0161]
[0162] where W′ k represents the updated weighted vector of the site; the hyperparameter η 1 = 0.0001, representing the learning rate of the loss function; W i represents the weighted weight vector set for each site, with a length of K; represents the Nabla operator, used to calculate the derivative; represents the mean squared error loss function; W kj represents the weight value of site j for target site k; Z j represents the predicted value of the output of the model of site j on the public sample; the hyperparameter ρ = 1, representing the control parameter, used to limit the weight; the optimizer selects Adam;
[0163] The calculated updated weighted vector W′ k , repeat K times in total, update the weight W′ corresponding to site k each time k , and finally obtain the updated weight set Aggregate the structure knowledge set obtained by the federated medical data structure knowledge extraction block to obtain the consensus structure knowledge A for site k t , and calculate A using the following formula t :
[0164]
[0165] Among them, denotes the structural knowledge extracted by site j, and the calculation formula is as follows:
[0166]
[0167] (4) The server samples new samples from the public medical dataset D 0 and broadcasts the public sample X 0 , the consensus structural knowledge A weighted by the weights extracted by the site collaboration enhancement module t , and the global model parameters θ to all K medical sites;
[0168] (5) For each medical site k participating in the training, update the local model θ k using the corresponding local private medical data; when the site receives the parameter θ sent by the server, first initialize the local model θ k = θ;
[0169] Use the Dice coefficient as the loss function, and the calculation formula is as follows:
[0170]
[0171] Among them, denotes the model output of the model parameters θ k for the sample X; y represents the label corresponding to the sample;
[0172] The medical site data satisfies the following conditions:
[0173] (X, y) ~ D k train
[0174] The size of the medical image is: H×W×D, where H represents the height of the medical image; W represents the width of the medical image; D represents the depth of the medical image; select H = W = D = 128;
[0175] Adopt the Adam optimizer to update the global model downloaded by the site; set the batch size of the training data to 1, and the learning rate η 2 = 0.0001, η 3 = 0.0001, and the loss function proportionality coefficient α = 0.5. Use the loss function defined by the following formula to update the segmentation model:
[0176]
[0177] Among them, Denote the loss function used in the current task; θ k Denote the local model of site k; θ′ k Denote the updated local model;
[0178] (6) Update the parameters of the local model for site k again:
[0179] Calculate the attention map of the output of the local model in the previous communication round on the client The attention map of the output of the model to be updated currently The set A of attention maps downloaded from the server t Update the model; calculate the updated objective function using the following formula
[0180]
[0181] For the medical image segmentation task, the Dice coefficient can be used as a metric, and the similarity metric function for the medical segmentation task is defined using the following formula:
[0182]
[0183] Utilize the loss function Update the parameters of the target model, select Adam as the optimizer, and the learning rate η 2 = 0.0001, η 3 = 0.0001, the proportionality coefficient α = 0.5, the dimensionality coefficient τ = 2. Combining the loss function of the previous updated model, it is expressed using the following formula:
[0184]
[0185] (7) When all medical sites have updated their models, send the updated local model θ′ k back to the server;
[0186] (8) The server aggregates all the local model parameters to obtain a new global model θ; where, p k denotes the weighting coefficient of site k, and the calculation formula is as follows:
[0187]
[0188] In the embodiment of the present invention, there are K = 8 medical sites. Therefore,
[0189] (9) The server samples new samples from the public medical dataset and simultaneously generates a set of attention maps using the new model sent by the client;
[0190] (10) Repeat the above steps (1) to (9) until the global model converges;
[0191] Define the metric for global model convergence: For all sites k ∈ K, the global model achieves the best segmentation performance on the test datasets of all sites, and it is considered that the global model converges; meanwhile, the Dice coefficient is used to measure the segmentation performance.
[0192] The experimental environment of the method of the present invention runs on a computer with a central processing unit (CPU) of Intel(R) Xeon(R) Gold 6230, a memory of 251GB, and a graphics card of Nvidia GeForce GTX 2080Ti; the local training of federated learning and the federated learning algorithm part are implemented based on the Python 3.9 programming language using Pytorch 1.10; the federated communication module is written using the RPC framework developed by Google.
[0193] The method of the present invention was experimentally tested using federated medical image data, specifically FeTS2021. Through data preprocessing and other means, a total of 8 medical sites were constructed. The data sampling devices used among the sites were different, and there were differences in the characteristics of the data at each medical site; the dataset selected was a labeled multi-modal multi-class dataset, and the sample features of the data included 4 sequences, namely T1, T2, T1CE, and magnetic resonance imaging fluid-attenuated inversion recovery sequence (Flair); these 4 sequences complement each other and can provide sufficient features for the segmentation of brain tumor images; there are a total of 3 mutually nested sub-regions in the data, namely the whole tumor (WT), the tumor core (TC) of the necrotic core and non-enhanced tumor core, and the enhancing tumor (ET).
[0194] The method proposed in the present invention was compared with other federated learning methods. The benchmark method was the well-known FedAvg. In addition, there were the federated integrated knowledge distillation methods FedDF and FedMD based on the same public dataset; and the federated learning methods FedBN and HarmoFL for dealing with the distribution differences of medical image data. In the experiment, all methods ran for 50 communication rounds, and the result with the best average performance in the 50 rounds was selected as the final result of the method; all methods were run 5 times and the average results were taken.
[0195] The final experimental results are shown in Table 1 below. It can be seen that the method of the present invention has achieved performance improvement in the segmentation results of the 3 sub-regions compared to the benchmark method; such as Figure 4Shown is a schematic diagram of the visualization demonstration results of the method of the present invention in an embodiment. It can be seen from the figure that the method of the present invention can extract an accurate attention map in the brain image and thus improve the accuracy of segmentation.
[0196] Method WT TC ET FedAvg 87.24±0.03 77.58±0.12 65.89±0.26 FedBN 85.51±0.06 73.07±0.14 66.46±0.24 HarmoFL 86.17±0.03 71.94±0.13 65.82±0.27 FedMD 85.82±0.03 76.82±0.12 67.83±0.22 FedDF 87.32±0.03 74.36±0.14 68.67±0.26 The method of the present invention 87.42±0.28 77.81±0.13 68.70±0.26
[0197] From the above experimental data, it can be seen that the method of the present invention uses a variety of deep learning models and federated learning algorithms on real medical data sets to verify the superiority of the method and demonstrate efficient segmentation performance; deep learning modeling in the medical field usually requires a large amount of real medical data sets for centralized training, and federated learning is a clever solution to the privacy and legal risks caused by centralized training models; in order to ensure the performance of the model established by federated learning, researchers must design federated learning models in combination with actual medical scenarios, and the method proposed in the present invention can help each medical site train and generate high-performance models when the local data differences are too large.
Claims
1. A method for sharing medical institution data based on federated learning, comprising the following steps: S1. Obtain a public medical dataset and establish an association between medical sites and service terminals participating in federated training; S2. Use the public medical dataset obtained in step S1 to construct a federated medical data structure knowledge extraction block and extract a set of structural knowledge; S3. Use the set of structural knowledge extracted in step S2 to construct a site collaboration enhancement module for dynamic weighting processing to obtain consensus structural knowledge; specifically including: (3-1) Define the model θ of site k k In the public sample X 0 The inference output is Z k = f(X 0 ; θ k ); Collect the inference outputs of all K sites, defined as (3-2) The weighted vector W of the updated site k is represented by the following formula k :[[-END]] Among them, W k ′ represents the weighted vector of the updated site; η 1 represents the learning rate of the loss function; W i represents the weighted weight vector set for each site, with a length of K; represents the Nabla operator, which is used to calculate the derivative; represents the mean squared error loss function; W kj represents the weight value of site j for target site k; Z j represents the predicted value of the model of site j at the output of the common sample; ρ represents the control parameter, which is used to limit the weight; (3-3) Repeat the weighted vector W of the updated site obtained by calculating in step (3-2) k ′, repeat a total of K times, and update the weight W corresponding to site k each time k ′, and finally obtain the updated weight set Aggregate the structural knowledge set extracted in step S2 to obtain the consensus structural knowledge A for site k t , and calculate A using the following formula t : Among them, represents the structural knowledge extracted by site j, and the calculation formula is as follows: where φ(X 0 ; f(X 0 ; θ j )) represents the structural knowledge module function extracted based on the sample X 0 and the model θ of site j j ; S4. Use the public medical dataset obtained in step S1 and the set of structural knowledge obtained in step S3 to construct a local update module and calculate the local overall loss; S5. Use the local overall loss calculated in step S4 to update the local model and complete the training process of the site; S6. Use the sites that have completed training in step S5 to perform aggregation processing on the server side to obtain a new global model of federated medical data; S7. Repeat the above steps S2 - S6, perform multiple rounds of federated training in a loop until the global model of the federated medical data obtained reaches the desired performance, and complete the construction of the federated medical data model; S8. Use the federated medical data model constructed in step S7 to achieve the sharing of medical institution data.
2. The method for sharing medical institution data based on federated learning according to claim 1, wherein the obtaining of the public medical dataset in step S1 and the establishment of an association between medical sites and service terminals participating in federated training specifically include: The public medical dataset includes a labeled dataset, an unlabeled dataset, and a third-party public dataset with differences in features and data volume from all participating federal training sites; define the public medical dataset as D 0 ; Establish a global model θ on the server side and initialize the training environment, including the initial parameters of the model and the initial values of the set of attention maps; define that there are K sites in the federated learning system, each site is represented as k, and the initialized local model of the site is Initialize the local model of the previous federated training as The server samples a new public sample X from the public medical dataset 0 ~D 0 , where X 0 represents the public sample.
3. The method for sharing medical institution data based on federated learning according to claim 2, wherein the using of the public medical dataset obtained in step S1 to construct a federated medical data structure knowledge extraction block and extract a set of structural knowledge in step S2 specifically include: (2-1) Site-local model θ k The predicted value output for the feature map The value of the average gradient on channel c is defined as β c , and β is calculated using the following formula c : Among them, H 1 represents the total longitudinal pixel length of the feature map M; W 1 represents the total transverse pixel length of the feature map M; i represents the i-th transverse pixel of the feature map M; j represents the j-th longitudinal pixel of the feature map M; y c represents the predicted value output by the local model θ k ; M c represents the value of the feature map in channel c; In a deep learning neural network driven by medical image data, select the bottom - most module of the encoder part in U - Net as the feature map M; (2 - 2) Aggregate the average value of the gradients calculated in step (2 - 1) with the feature map, filter the unresponsive parts, and obtain structural knowledge. The calculation formula is as follows: Among them, A c represents the value of the structural knowledge in channel c; C represents the total number of channels; c represents the channel with the largest labeled data volume; ReLU(·) represents the non-linear activation function, and the calculation formula is as follows: Use the obtained structural knowledge to replace the parameters of the model as the medium for knowledge transfer between sites, and reduce the inconsistency of semantic knowledge captured by models at different sites; For site k ∈ K, define the corresponding output as Use the following defined function to represent the structural knowledge module in steps (2 - 1) and (2 - 2): φ(X,f(X;θ)) where X represents the sample input to the function; f(·) represents the model used in the inference process; θ represents the parameters corresponding to the model; Adopt the common sample X 0 , the local models of all K sites At the server side, through the constructed knowledge structure module φ(X, f(X; θ)), extract the structural knowledge set where φ(X 0 ; f(X 0 ; θ k )) represents the structural knowledge module function extracted based on the sample X 0 and the model θ of site k k .
4. The method for sharing medical institution data based on federated learning according to claim 3, wherein the using of the public medical dataset obtained in step S1 and the set of structural knowledge obtained in step S3 to construct a local update module and calculate the local overall loss in step S4 specifically include: (4-1) For each medical site k participating in the training, download the global model θ and the consensus structure knowledge A from the server t and the new public samples X 0 to the local site, and the site uses the samples and the consensus structure knowledge to align the feature space of the site model; (4 - 2) Site k updates the parameters of the local model: The local loss function of the site is expressed by the following formula where α represents a proportionality coefficient; represents the task loss function; represents the local integrated knowledge distillation loss function.
5. The method for sharing medical institution data based on federated learning according to claim 4, wherein the task loss function included in the local loss function specifically includes: Adopt the local dataset D k ={(X i , y i )}, define the task loss function as follows: Among them, X i represents the i-th batch of samples in the local data; y i represents the label corresponding to the i-th batch of samples in the local data; f(X i ; θ k ) represents the output of the local model; loss represents a function, defined as: In the medical segmentation task, represents the sample space of the data, and the corresponding label space is Among them, C represents the number of channels of the sample image; H represents the height of the sample image; W represents the width of the sample image; Calculate the task loss function using the following formula: Among them, y represents the label corresponding to the sample; represents the predicted result, and the calculation formula is as follows:
6. The method for sharing medical institution data based on federated learning according to claim 5, characterized in that the local structural knowledge integration loss included in the local loss function specifically includes: Indicates the similarity between the attention map extracted by the local model and the mean of the set of attention maps downloaded from the server. Indicates the similarity between the attention map extracted by the local model and the attention map extracted by the local model in the previous communication round. The local structure knowledge integration loss is calculated using the following formula Among them, represents the local model θ to be updated currently k The output structural knowledge is calculated as follows: Among them, φ(X 0 ; f(X 0 ; θ i )) represents the structural knowledge module function extracted based on the sample X 0 and the model θ i ; A t Represents the consensus structure knowledge obtained by downloading through the server; Represents the local model obtained based on the previous round of federated training The output structural knowledge, and the calculation formula is as follows: Among them, represents the structural knowledge module function extracted based on the sample X 0 and the model ; τ represents the temperature coefficient, which is used to control the scope covered by the structural knowledge; 0≤sim(p,q)≤1 represents the similarity metric function of inputs p and q, and the closer the value of the function is to 1, the more similar p and q are; For medical data, the Dice coefficient is used to implement the similarity function, and the calculation formula is as follows:
7. The method for sharing medical institution data based on federated learning according to claim 6, characterized in that using the local overall loss calculated in step S4 to update the local model and complete the training process of the site specifically includes: For any site k, the update process is represented by the following formula: Among them, θ k represents the old local model of site k; θ k ' represents the updated local model of site k; η 2 represents the task loss learning rate; η 3 represents the learning rate of the local integrated knowledge distillation loss function ; α represents the proportionality coefficient of the two losses; represents the gradient used to update the local model; For any site k, store the currently updated federated model and use it for the next parameter update, which is represented by the following formula: After the training of site k is completed, the updated local model θ k ′ is assigned as the local model stored by site k in the previous communication round; then, the client submits the updated local model θ k ′ to the server; the server waits for all K sites to submit their local models to obtain a model set 8. The method for sharing medical institution data based on federated learning according to claim 7, characterized in that the sites that have completed training in step S6 are aggregated on the server side to obtain a new global model of federated medical data, specifically including: Aggregate the local model parameter sets submitted by all K sites on the server side Obtain a new global model of federated medical data where p k represents the weighted coefficient of the site, and the calculation formula is as follows:
Citation Information
Patent Citations
Heterogeneous model aggregation method and system based on federated learning
CN113705610A
Method, device, and apparatus for combining horizontal federation and vertical federation, and medium
WO2021083276A1