A model training method and device based on specific federated learning

By adopting a model training method based on specific federated learning in magnetic resonance imaging, the global shared model and local model are separated, and the weighted comparison regularization loss function is used to solve the domain shift and privacy leakage problems of federated learning in MR image reconstruction, achieving high-precision image reconstruction.

CN114627202BActive Publication Date: 2025-05-16HARBIN INST OF TECH SHENZHEN GRADUATE SCHOOL
View PDF 2 Cites 0 Cited by

Patent Information

Application Number
CN202210212867.8
Authority / Receiving Office
CN · China
Patent Type
Patents(China)
Current Assignee / Owner
Filing Date
2022-03-04
Publication Date
2025-05-16
Estimated Expiration
2042-03-04

AI Technical Summary

Technical Problem

In magnetic resonance imaging, deep learning-based image reconstruction methods are difficult to apply because they require a large number of diverse pairing data, and the application of federated learning in MR image reconstruction has problems of domain shifting and privacy leakage.

Method used

A model training method based on specific federated learning is proposed. By separating the global shared model and local model between the server and the client, the weighted comparison regularization loss function is used to correct the update direction of global generalization, reducing domain drift and meeting the privacy protection mechanism.

Benefits of technology

While meeting the privacy protection mechanism, it alleviates the domain drift of the client during the training process, promotes model convergence, and improves the accuracy of MR image reconstruction.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN114627202B_ABST
    Figure CN114627202B_ABST
Patent Text Reader

Abstract

The present application provides a model training method and device based on specific federated learning, the method comprising: in each round of communication, the server sends the global shared model to each client, each client performs a local gradient update according to the global shared model currently transmitted by the server, after the local update is completed, the client participates in the global gradient update of the server, and returns the update result to the server, the server determines the next round of global shared model according to the update result returned by the client, and introduces weighted contrast regularization from the second round to correct the local gradient update of the client; after multiple rounds of communication, the client gradually has the characteristics of the global shared model. The present application can alleviate the domain drift of the client during the training process while satisfying the privacy protection mechanism and promote convergence.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present application relates to the field of image processing technology, and in particular to a model training method and device based on specific federated learning. Background Art

[0002] Magnetic resonance (MR) imaging has become a mainstream diagnostic tool in radiology and medicine. However, its complex imaging process results in longer acquisition times than other methods such as computed tomography (CT), x-rays, and ultrasound. In order to reduce scanning time and improve patient experience, several methods have been proposed to speed up MRI, such as traditional methods based on compressed sensing, dictionary learning, low rank, etc. In recent years, data-driven deep learning methods have also achieved significant improvements in MR image reconstruction, mainly due to the large amount of available training data. However, the superior results obtained by deep learning-based methods often rely on a large amount of diverse paired data, which is difficult to collect due to patient privacy issues.

[0003] Recently, Federated Learning (FL) algorithms have been proposed, which provide a platform for different clients to collaboratively learn using local computing power, memory, and data without sharing any private local data. FedAvg is one of the standard and most widely used FL algorithms, which collects local models from each client in each round of communication and distributes their average to each client for the next update. Due to the distributed federated training, FL has been applied in many fields, including image classification, object detection, domain generalization, medical image segmentation, etc. However, in MR image reconstruction, there is heterogeneity in different MRI scanners and imaging protocols in different hospitals, resulting in domain shifts between clients. Unfortunately, under these conditions, simple federated training using FL-trained models can still be suboptimal. Technologists have tried to solve this problem by repeatedly adjusting and aligning latent features between source and target clients, which is the first attempt to use FL in MR image reconstruction.

[0004] Although FL has been applied to MR image reconstruction, this cross-site approach often requires sacrificing one client as the target location in order to align with other clients in each round of communication. Obviously, any client used as a target site will lead to privacy leakage issues, and the cross-site approach contradicts the purpose of FL, which is to prevent clients from communicating with each other through local data. In addition, when the number of clients is large, the process becomes cumbersome due to repeated training and frequent feature exchanges. More importantly, this mechanism can only learn a general global model while ignoring the specific properties of each client. Previous research on domain adaptation has also shown that encoders are often used to learn shared representations to ensure that all inputs are equally suitable for any domain transformation. Therefore, although the FL algorithm has made initial attempts in MR image reconstruction, its accuracy still needs to be improved. Summary of the invention

[0005] In view of the above problems, the present application is proposed to provide a model training method and device based on specific federated learning to overcome the above problems or at least partially solve the above problems, including:

[0006] A model training method based on specific federated learning is used in a machine learning system, wherein the machine learning system includes a server and at least two clients; the server stores a global shared model, a set of global shared models completed in the previous round of training, and training data, and for the first round of training, the set of global shared models completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively; the model training method is for the server; the model training method includes:

[0007] The server sends the global shared model to each of the clients; the client is used to receive the global shared model sent by the server; the local model is trained according to the global shared model and the local data to obtain a trained local model; the global shared model is trained according to the trained local model and the local data to obtain a preliminarily trained global shared model; the preliminarily trained global shared model is sent to the server;

[0008] The server receives the global shared model after preliminary training sent by each of the clients;

[0009] When the global shared model set completed in the previous round of training is not empty, the server trains each of the global shared models completed in the preliminary training according to the global shared model, the global shared model set completed in the previous round of training and the training data to obtain a global shared model set completed in the training;

[0010] The server updates the global shared model according to the trained global shared model set.

[0011] Preferably, after the step of receiving the global shared model that has been preliminarily trained and sent by each of the clients, the server further includes:

[0012] When the set of global shared models that have completed the previous round of training is empty, the server trains each of the global shared models that have completed the preliminary training according to the training data to obtain a set of global shared models that have completed the training;

[0013] The server updates the global shared model according to the trained global shared model set.

[0014] Preferably, the trained global shared model set includes all trained global shared models; the server side trains each of the preliminarily trained global shared models based on the global shared model, the set of global shared models trained in the previous round, and the training data to obtain the trained global shared model set, comprising:

[0015] For each global shared model that has completed preliminary training, the server performs the following steps:

[0016] The server processes the training data according to the global shared model to obtain a first prediction result;

[0017] The server processes the training data according to the global shared model set completed in the previous round of training to obtain a second prediction result set;

[0018] The server processes the training data according to the global shared model that has been preliminarily trained to obtain a third prediction result;

[0019] The server determines a first loss value according to the first prediction result, the second prediction result set, the third prediction result and a pre-constructed weighted comparison regularization loss function;

[0020] The server determines a second loss value according to the third prediction result and a pre-constructed supervised reconstruction loss function;

[0021] The server side trains the preliminarily trained global shared model according to the first loss value and the second loss value to obtain the trained global shared model.

[0022] Preferably, the trained global shared model set includes all trained global shared models; and the step of updating the global shared model on the server side according to the trained global shared model set includes:

[0023] The server side sets the average value of all the trained global shared models as the global shared model.

[0024] A model training method based on specific federated learning is used in a machine learning system, the machine learning system comprising a server and at least two clients; the server stores a global shared model, a set of global shared models completed in the previous round of training, and training data, and for the first round of training, the set of global shared models completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively; the model training method is for any one of the at least two clients; the model training method comprises:

[0025] The client receives the global shared model sent by the server;

[0026] The client trains the local model according to the global shared model and the local data to obtain a trained local model;

[0027] The client trains the global shared model according to the trained local model and the local data to obtain a preliminarily trained global shared model;

[0028] The client sends the global shared model that has completed preliminary training to the server; the server is used to receive the global shared model that has completed preliminary training sent by each of the clients; when the set of global shared models that have completed the previous round of training is not empty, each global shared model that has completed preliminary training is trained according to the global shared model, the set of global shared models that have completed the previous round of training and the training data to obtain a set of trained global shared models; based on the set of trained global shared models, the global shared model set is updated.

[0029] Preferably, the client trains the local model based on the global shared model and the local data to obtain the trained local model, comprising:

[0030] The client processes the local data according to the global shared model to obtain a fourth prediction result;

[0031] The client processes the local data according to the local model to obtain a fifth prediction result;

[0032] The client determines a third loss value according to the fourth prediction result, the fifth prediction result and a pre-constructed local loss function;

[0033] The client trains the local model according to the third loss value to obtain the trained local model.

[0034] Preferably, the client trains the global shared model based on the trained local model and the local data to obtain a preliminarily trained global shared model, including:

[0035] The client processes the local data according to the trained local model to obtain a sixth prediction result;

[0036] The client processes the local data according to the global shared model to obtain a seventh prediction result;

[0037] The client determines a fourth loss value according to the sixth prediction result, the seventh prediction result and a pre-constructed shared loss function;

[0038] The client trains the global shared model according to the fourth loss value to obtain the trained global shared model.

[0039] A model training device based on specific federated learning is used in a machine learning system, the machine learning system comprising a server and at least two clients; the server stores a global shared model, a set of global shared models completed in the previous round of training, and training data, and for the first round of training, the set of global shared models completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively; the model training device is for the server; the model training device comprises:

[0040] A global shared model sending module, used to send the global shared model to each of the clients; the client is used to receive the global shared model sent by the server; the local model is trained according to the global shared model and the local data to obtain a trained local model; the global shared model is trained according to the trained local model and the local data to obtain a preliminarily trained global shared model; the preliminarily trained global shared model is sent to the server;

[0041] A primary model receiving module, used for receiving the global shared model after preliminary training sent by each of the clients;

[0042] A primary model training module, used for training each of the global shared models that have been preliminarily trained according to the global shared model, the global shared model set that has been trained in the previous round and the training data to obtain a global shared model set that has been trained, when the global shared model set that has been trained in the previous round is not empty;

[0043] The global model determination module is used to update the global shared model according to the trained global shared model set.

[0044] A model training device based on specific federated learning is used in a machine learning system, the machine learning system comprising a server and at least two clients; the server stores a global shared model, a set of global shared models completed in the previous round of training, and training data, and for the first round of training, the set of global shared models completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively; the model training device is for any one of the at least two clients; the model training device comprises:

[0045] A global shared model receiving module, used for receiving the global shared model sent by the server;

[0046] A local model training module, used to train the local model according to the global shared model and the local data to obtain a trained local model;

[0047] A global shared model training module, used to train the global shared model based on the trained local model and the local data to obtain a global shared model that has been preliminarily trained;

[0048] A primary model sending module is used to send the global shared model that has completed preliminary training to the server end; the server end is used to receive the global shared model that has completed preliminary training sent by each of the clients; when the set of global shared models that have completed the previous round of training is not empty, each of the global shared models that have completed preliminary training is trained according to the global shared model, the set of global shared models that have completed the previous round of training and the training data to obtain a set of trained global shared models; based on the set of trained global shared models, the global shared model set is updated.

[0049] A machine learning system comprises a server and at least two clients; the server stores a global shared model, a set of global shared models completed in a previous round of training, and training data, and for a first round of training, the set of global shared models completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively;

[0050] The server is used to send the global shared model to each of the clients;

[0051] The client is used to receive the global shared model sent by the server;

[0052] The client is further used to train the local model according to the global shared model and the local data to obtain a trained local model;

[0053] The client is further used to train the global shared model according to the trained local model and the local data to obtain a preliminarily trained global shared model;

[0054] The client is also used to send the global shared model that has been initially trained to the server;

[0055] The server is further configured to receive the global shared model after preliminary training sent by each of the clients;

[0056] The server is further configured to, when the set of global shared models trained in the previous round is not empty, train each of the global shared models trained in the previous round according to the global shared model, the set of global shared models trained in the previous round and the training data to obtain a set of global shared models trained;

[0057] The server side is also used to update the global shared model based on the trained global shared model set.

[0058] This application has the following advantages:

[0059] In an embodiment of the present application, the global shared model is sent to each of the clients through the server; the client is used to receive the global shared model sent by the server; the local model is trained according to the global shared model and the local data to obtain a trained local model; the global shared model is trained according to the trained local model and the local data to obtain a global shared model that has been preliminarily trained; the preliminarily trained global shared model is sent to the server; the server receives the preliminarily trained global shared model sent by each of the clients; when the set of global shared models trained in the previous round is not empty, the server trains each of the preliminarily trained global shared models according to the global shared model, the set of global shared models trained in the previous round and the training data to obtain a set of trained global shared models; the server updates the global shared model according to the set of trained global shared models, which can alleviate the domain drift of the client during the training process while satisfying the privacy protection mechanism and promote convergence. BRIEF DESCRIPTION OF THE DRAWINGS

[0060] In order to more clearly illustrate the technical solution of the present application, the drawings required for use in the description of the present application will be briefly introduced below. Obviously, the drawings described below are only some embodiments of the present application. For ordinary technicians in this field, other drawings can be obtained based on these drawings without paying any creative work.

[0061] Figure 1 This is a schematic diagram of a framework overview of a model training method based on specific federated learning provided in one embodiment of the present application;

[0062] Figure 2 It is a flowchart of the steps of a model training method based on specific federated learning provided in one embodiment of the present application;

[0063] Figure 3 It is a flowchart of the steps of a model training method based on specific federated learning provided in one embodiment of the present application;

[0064] Figure 4 It is a flowchart of the steps of a model training method based on specific federated learning provided in one embodiment of the present application;

[0065] Figure 5 It is a visualization diagram of a potential feature T-SNE based on fastMRI, Brats, SMS and uMR data sets provided in one embodiment of the present application.

[0066] Figure 6It is a structural block diagram of a model training device based on specific federated learning provided in one embodiment of the present application;

[0067] Figure 7 It is a structural block diagram of a model training device based on specific federated learning provided in one embodiment of the present application;

[0068] Figure 8 It is a structural diagram of a computer device provided in one embodiment of the present application.

[0069] The reference numerals in the drawings of the specification are as follows:

[0070] 12. Computer equipment; 14. External devices; 16. Processing unit; 18. Bus; 20. Network adapter; 22. I / O interface; 24. Display; 28. Memory; 30. Random access memory; 32. Cache memory; 34. Storage system; 40. Program / utility; 42. Program module. DETAILED DESCRIPTION

[0071] In order to make the objects, features and advantages of the present application more obvious and understandable, the present application is further described in detail below in conjunction with the accompanying drawings and specific implementation methods. Obviously, the described embodiments are part of the embodiments of the present application, rather than all of the embodiments. Based on the embodiments in the present application, all other embodiments obtained by ordinary technicians in the field without creative work are within the scope of protection of the present application.

[0072] Reference Figure 1 In order to address the problem of low accuracy of multi-institutional federated reconstruction under the influence of domain drift, this application divides the MR image reconstruction model into two parts: a global shared model stored on the server side for learning generalized representations, and a local model stored on the client side for exploring the uniqueness of the client's domain distribution. In addition, in order to reduce the offset between the server and the client, this application also introduces a weighted contrast regularization function to correct the update direction of global generalization. Specifically, the global shared model (anchor point) preliminarily trained by the client is pulled toward the global shared model (positive point), and pushed away from the set of global shared models completed in the previous round of training (negative point). This application can alleviate the domain drift of the client during the training process while satisfying the privacy protection mechanism, promote convergence, and achieve significant improvement in model performance.

[0073] Reference Figure 2, shows a model training method based on specific federated learning provided by an embodiment of the present application, the model training method is used in a machine learning system, the machine learning system includes a server and at least two clients; the server stores a global shared model, a set of global shared models completed in the previous round of training, and training data, and for the first round of training, the set of global shared models completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively; the model training method is for the server; the model training method includes:

[0074] S110, the server sends the global shared model to each of the clients; the client is used to receive the global shared model sent by the server; the local model is trained according to the global shared model and the local data to obtain a trained local model; the global shared model is trained according to the trained local model and the local data to obtain a preliminarily trained global shared model; the preliminarily trained global shared model is sent to the server;

[0075] S120, the server receives the global shared model that has been initially trained and sent by each of the clients;

[0076] S130, when the set of global shared models trained in the previous round is not empty, the server trains each of the global shared models trained in the preliminary way according to the global shared model, the set of global shared models trained in the previous round and the training data to obtain a set of global shared models trained;

[0077] S140: The server updates the global shared model according to the trained global shared model set.

[0078] In an embodiment of the present application, the global shared model is sent to each of the clients through the server; the client is used to receive the global shared model sent by the server; the local model is trained according to the global shared model and the local data to obtain a trained local model; the global shared model is trained according to the trained local model and the local data to obtain a global shared model that has been preliminarily trained; the preliminarily trained global shared model is sent to the server; the server receives the preliminarily trained global shared model sent by each of the clients; when the set of global shared models trained in the previous round is not empty, the server trains each of the preliminarily trained global shared models according to the global shared model, the set of global shared models trained in the previous round and the training data to obtain a set of trained global shared models; the server updates the global shared model according to the set of trained global shared models, which can alleviate the domain drift of the client during the training process while satisfying the privacy protection mechanism and promote convergence.

[0079] Below, a model training method based on specific federated learning in this exemplary embodiment will be further described.

[0080] As described in step S110, the server sends the global shared model to each of the clients; the client is used to receive the global shared model sent by the server; the local model is trained according to the global shared model and the local data to obtain a trained local model; the global shared model is trained according to the trained local model and the local data to obtain a global shared model that has been preliminarily trained; and the preliminarily trained global shared model is sent to the server.

[0081] It should be noted that the server side can cyclically execute the training steps as described in S110-S140, that is, the global shared model obtained by the output of each round of training is used as the global shared model for the input of the next round of training. The global shared model stored on the server side and the set of global shared models completed in the previous round of training, as well as the local model stored on the client are updated during each round of training. The model involved in this application can be a neural network model, such as a convolutional neural network model, a recurrent neural network model, a deep residual network model, etc. This application does not limit the specific category of the model involved.

[0082] The server sends the global shared model to each of the clients, which can be understood as the server sending the complete global shared model, or as the server sending all weight parameters of the global shared model, or as the server sending part of the weight parameters of the global shared model, where the part of the weight parameters refers to the global shared model having updated weight parameters compared to the previous round of global shared model.

[0083] As described in step S120, the server receives the global shared model that has been initially trained and sent by each of the clients.

[0084] Since different clients have different data processing speeds, the server may execute step S130 after receiving the global shared models that have been preliminarily trained sent by all the clients, or may process the global shared models that have been preliminarily trained sent by each client separately in the order of receipt, and execute step S140 after processing all the global shared models that have been preliminarily trained.

[0085] The server side receives the global shared model that has been initially trained and sent by each of the clients. It can be understood that the server side receives the complete global shared model that has been initially trained, or it can be understood that the server side receives all weight parameters of the global shared model that has been initially trained, or it can be understood that the server side receives part of the weight parameters of the global shared model that has been initially trained, and the part of the weight parameters refers to the global shared model that has been initially trained has updated weight parameters compared to the global shared model.

[0086] As described in step S130, when the set of global shared models completed in the previous round of training is not empty, the server side trains each global shared model that has completed the preliminary training based on the global shared model, the set of global shared models completed in the previous round of training and the training data to obtain a set of global shared models that have completed the training.

[0087] When the set of global shared models completed in the previous round of training is not empty, that is, starting from the second round of training, the server side trains each of the global shared models completed through the preliminary training through the training data, the pre-constructed supervised reconstruction loss function and the weighted contrast regularization loss function, and obtains a trained global shared model corresponding to each of the global shared models completed through the preliminary training, forming the trained global shared model set including all the trained global shared models.

[0088] Since each round of training needs to be performed alternately between client updates and server updates, the MR image reconstruction model is considered to be divided into the global shared model stored on the server and the local model stored on the kth client to share global information and find unique depth information. The supervised reconstruction loss function can be expressed as:

[0089]

[0090] Among them, G e and Respectively represent the global shared model and the local model, x∈C M Represents an undersampled image, y represents a fully sampled image, x and y constitute the training data pre-stored in the server, and K represents the total number of clients. It should be noted that the global shared model is learned jointly by the server and the client. Although the client has shared the global shared model that has been preliminarily trained with the server to find a common representation between multiple clients, there is always an offset between the global shared model and the global shared model that has been preliminarily trained during the iterative optimization process, which is mainly caused by the domain shift during the local optimization process. In order to further correct the local update and make the model have global recognition capabilities, the present application introduces weighted contrast regularization between the global shared model and the local model, forcing the global shared model to learn a stronger generalized representation. Unlike traditional contrastive learning, the present application does not need to find positive and negative pairs from the data, but directly regularizes the update direction of the network parameters. This allows gradient updates to be corrected more directly without relying on a large number of training samples in each iteration.

[0091] Assuming that the kth client performs a local update, it first receives the global shared model from the server, and then performs a local iterative update based on this data. However, the global parameters from the server always have a smaller deviation than the local parameters. The application defines the weighted contrast regularization loss function as:

[0092]

[0093] Combined with the supervised reconstruction loss function, the overall loss function of the model can be expressed as:

[0094]

[0095] Among them, μ is a hyperparameter that controls the weight of the weighted contrast regularization loss function.

[0096] After obtaining the trained global shared model set, the server side also updates the global shared model set trained in the previous round based on the trained global shared model set to ensure that the global shared model set trained in the previous round stored on the server side before the start of the next round of communication is the trained global shared model set obtained during this round of communication.

[0097] As described in step S140, the server updates the global shared model based on the trained global shared model set.

[0098] The server side may use a variety of fusion algorithms to fuse multiple trained global shared models included in the trained global shared model set to update the global shared model. For example, the multiple trained global shared models may be averaged to update the global shared model, or the multiple trained global shared models may be weighted to update the global shared model, or other preset algorithms may be used to process multiple trained global shared models to update the global shared model.

[0099] Reference Figure 3 In one embodiment of the present application, after step S120, the following steps are further included:

[0100] S210, when the set of global shared models trained in the previous round is empty, the server trains each of the global shared models trained initially according to the training data to obtain a set of global shared models trained;

[0101] S220: The server updates the global shared model according to the trained global shared model set.

[0102] As described in step S210, when the set of global shared models that have completed the previous round of training is empty, the server side trains each of the global shared models that have completed the preliminary training according to the training data to obtain a set of global shared models that have completed the training.

[0103] When the set of global shared models completed in the previous round of training is empty, that is, it is the first round of training, the server side trains each of the global shared models completed through the preliminary training through the training data and the supervised reconstruction loss function, obtains a trained global shared model corresponding to each of the global shared models completed through the preliminary training, and forms the trained global shared model set including all the trained global shared models.

[0104] As described in step S220, the server updates the global shared model based on the trained global shared model set.

[0105] The server side may use a variety of fusion algorithms to fuse multiple trained global shared models included in the trained global shared model set to update the global shared model. For example, the multiple trained global shared models may be averaged to update the global shared model, or the multiple trained global shared models may be weighted to update the global shared model, or other preset algorithms may be used to process multiple trained global shared models to update the global shared model.

[0106] In this embodiment, the trained global shared model set includes all trained global shared models; the server trains each of the preliminarily trained global shared models according to the global shared model, the set of global shared models trained in the previous round, and the training data, and the step of obtaining the trained global shared model set includes:

[0107] For each global shared model that has completed preliminary training, the server performs the following steps:

[0108] The server processes the training data according to the global shared model to obtain a first prediction result;

[0109] The server processes the training data according to the global shared model set completed in the previous round of training to obtain a second prediction result set;

[0110] The server processes the training data according to the global shared model that has been preliminarily trained to obtain a third prediction result;

[0111] The server determines a first loss value according to the first prediction result, the second prediction result set, the third prediction result and a pre-constructed weighted comparison regularization loss function;

[0112] The server determines a second loss value according to the third prediction result and a pre-constructed supervised reconstruction loss function;

[0113] The server side trains the preliminarily trained global shared model according to the first loss value and the second loss value to obtain the trained global shared model.

[0114] Specifically, the server determines a first overall loss value based on the first loss value and the second loss value, and trains the global shared model that has completed the preliminary training based on the first overall loss value, and stops training when the first overall loss value is less than a first preset value, thereby obtaining the trained global shared model.

[0115] Reference Figure 4 , shows a model training method based on specific federated learning provided by an embodiment of the present application, which is used in a machine learning system, wherein the machine learning system includes a server and at least two clients; the server stores a global shared model, a set of global shared models completed in the previous round of training, and training data, and for the first round of training, the set of global shared models completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively; the model training method is for any one of the at least two clients; the model training method includes:

[0116] S310, the client receives the global shared model sent by the server;

[0117] S320, the client trains the local model according to the global shared model and the local data to obtain a trained local model;

[0118] S330, the client trains the global shared model according to the trained local model and the local data to obtain a preliminarily trained global shared model;

[0119] S340, the client sends the global shared model that has completed preliminary training to the server; the server is used to receive the global shared model that has completed preliminary training sent by each client; when the set of global shared models that have completed the previous round of training is not empty, each global shared model that has completed preliminary training is trained according to the global shared model, the set of global shared models that have completed the previous round of training and the training data to obtain a set of trained global shared models; based on the set of trained global shared models, the global shared model set, the global shared model is updated.

[0120] As described in step S310, the client receives the global shared model sent by the server.

[0121] The client receives the global shared model sent by the server, which can be understood as the client receiving the complete global shared model, or as the client receiving all weight parameters of the global shared model, or as the client receiving part of the weight parameters of the global shared model, where the part of the weight parameters refers to the global shared model having updated weight parameters compared to the previous round of global shared model.

[0122] As described in step S320, the client trains the local model based on the global shared model and the local data to obtain a trained local model.

[0123] The client updates the local gradient according to the global shared model sent by the server to find the optimal unique local information, as shown below:

[0124]

[0125] Among them, L cl is a pre-built local loss function,

[0126] The advantage of this update rule is that the number of local updates can be controlled, so as to find the optimal trained local model specific to the client based on local data.

[0127] After obtaining the trained local model, the client also updates the local model based on the trained local model to ensure that the local model stored in the client before the next round of communication starts is the trained local model obtained during this round of communication.

[0128] As described in step S330, the client trains the global shared model based on the trained local model and the local data to obtain a preliminarily trained global shared model.

[0129] After the local update is completed, the client participates in the global gradient update as follows:

[0130]

[0131] Among them, L se is a pre-built shared loss function.

[0132] As described in step S340, the client sends the global shared model that has completed preliminary training to the server; the server is used to receive the global shared model that has completed preliminary training sent by each client; when the set of global shared models that have completed the previous round of training is not empty, each global shared model that has completed preliminary training is trained according to the global shared model, the set of global shared models that have completed the previous round of training and the training data to obtain a set of trained global shared models; based on the set of trained global shared models, the global shared model set is updated.

[0133] The client sends the global shared model that has completed the preliminary training to the server. It can be understood that the client sends the complete global shared model that has completed the preliminary training, or it can be understood that the client sends all weight parameters of the global shared model that has completed the preliminary training, or it can be understood that the client sends part of the weight parameters of the global shared model that has completed the preliminary training, and the part of the weight parameters refers to the global shared model that has completed the preliminary training has updated weight parameters compared to the global shared model.

[0134] In this embodiment, the client trains the local model according to the global shared model and the local data to obtain the trained local model, including:

[0135] The client processes the local data according to the global shared model to obtain a fourth prediction result;

[0136] The client processes the local data according to the local model to obtain a fifth prediction result;

[0137] The client determines a third loss value according to the fourth prediction result, the fifth prediction result and a pre-constructed local loss function;

[0138] The client trains the local model according to the third loss value to obtain the trained local model.

[0139] Specifically, the client trains the local model according to the third loss value, stops training when the third loss value is less than a second preset value, and obtains the trained local model.

[0140] In this embodiment, the client trains the global shared model based on the trained local model and the local data to obtain a preliminary trained global shared model, including:

[0141] The client processes the local data according to the trained local model to obtain a sixth prediction result;

[0142] The client processes the local data according to the global shared model to obtain a seventh prediction result;

[0143] The client determines a fourth loss value according to the sixth prediction result, the seventh prediction result and a pre-constructed shared loss function;

[0144] The client trains the global shared model according to the fourth loss value to obtain the trained global shared model.

[0145] Specifically, the client trains the global shared model according to the fourth loss value, and stops training when the fourth loss value is less than a third preset value, thereby obtaining the trained global shared model.

[0146] The specific steps of this application are shown in the following algorithm. In each round of communication, the server sends the global shared model to each of the clients. Then, each client performs a local gradient update based on the global shared model to obtain its optimal unique information, as shown in formula (2). Then, the client participates in the update of the server according to formula (3), and then the server corrects the local gradient update according to formula (5).

[0147] Input: K client data: D1, D2, ..., D k ; The number of client updates T; The number of communication rounds Z; The hyperparameter μ; The learning rate η corresponding to each client k ;

[0148] Output: Global shared model

[0149]

[0150] Algorithm 1

[0151] Reference Figure 5 To verify the performance of the model training method provided in this application, the T-SNE distribution of the potential features is visualized, where (ad) shows the SingleSet, FedAvg, excluding L conThe FedMRI algorithm of FL and the algorithm of this application. In SingleSet, each client is trained using only their local data. The distribution of points in (a) is obviously different because each dataset has its own bias, while the data in (b), (c), and (d) overlap to varying degrees because these models benefit from the federated training mechanism of FL. However, for datasets with large distribution differences, such as fastMRI and BraTS, FedAvg almost fails (see Figure 5 (b)).

[0152] It is worth noting that even without L con , the method of this application can still align the latent space distributions on four different datasets, which shows that sharing a global shared model and maintaining a client-specific local model can effectively reduce the domain shift problem (see Figure 5 (c)). Figure 5 (d) shows that the latent feature distributions of different clients are clearly fully mixed. This can be attributed to the fact that the weighted contrast regularization enables the algorithm of the present application to effectively correct the deviation between the client and the server during optimization (see Figure 5 (d)).

[0153] In one embodiment of the present application, there is also provided an image processing method based on specific federated learning, which is used in a machine learning system. The machine learning system includes a server and at least two clients. The image reconstruction method is for the server. The image processing method includes:

[0154] Get the data to be processed;

[0155] The data to be processed is processed according to the global shared model trained based on any of the above-mentioned model training methods to obtain an image reconstruction result of the data to be processed.

[0156] As for the device embodiment, since it is basically similar to the method embodiment, the description is relatively simple, and the relevant parts can be referred to the partial description of the method embodiment.

[0157] Reference Figure 6 , shows a model training device based on specific federated learning provided by an embodiment of the present application, which is used in a machine learning system, wherein the machine learning system includes a server and at least two clients; the server stores a global shared model, a set of global shared models completed in the previous round of training, and training data. For the first round of training, the set of global shared models completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively; the model training device is for the server; the model training device includes:

[0158] The global shared model sending module 410 is used to send the global shared model to each of the clients; the client is used to receive the global shared model sent by the server; the local model is trained according to the global shared model and the local data to obtain a trained local model; the global shared model is trained according to the trained local model and the local data to obtain a preliminarily trained global shared model; the preliminarily trained global shared model is sent to the server;

[0159] The primary model receiving module 420 is used to receive the global shared model after preliminary training sent by each client;

[0160] The primary model training module 430 is used to train each of the global shared models that have been preliminarily trained according to the global shared model, the global shared model set that has been trained in the previous round, and the training data to obtain a global shared model set that has been trained when the global shared model set that has been trained in the previous round is not empty;

[0161] The global model determination module 440 is used to update the global shared model according to the trained global shared model set.

[0162] Reference Figure 7 , shows a model training device based on specific federated learning provided by an embodiment of the present application, which is used in a machine learning system, wherein the machine learning system includes a server and at least two clients; the server stores a global shared model, a set of global shared models completed in the previous round of training, and training data, and for the first round of training, the set of global shared models completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively; the model training device is for any one of the at least two clients; the model training device includes:

[0163] A global shared model receiving module 510 is used to receive the global shared model sent by the server;

[0164] A local model training module 520 is used to train the local model according to the global shared model and the local data to obtain a trained local model;

[0165] A global shared model training module 530 is used to train the global shared model according to the trained local model and the local data to obtain a preliminarily trained global shared model;

[0166] The primary model sending module 540 is used to send the global shared model that has completed preliminary training to the server end; the server end is used to receive the global shared model that has completed preliminary training sent by each of the clients; when the set of global shared models that have completed the previous round of training is not empty, each of the global shared models that have completed preliminary training is trained according to the global shared model, the set of global shared models that have completed the previous round of training and the training data to obtain a set of trained global shared models; based on the set of trained global shared models, the global shared model set is updated.

[0167] In one embodiment of the present application, a machine learning system is further provided, comprising a server and at least two clients; the server stores a global shared model, a set of global shared models completed in the previous round of training, and training data, and for the first round of training, the set of global shared models completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively;

[0168] The server is used to send the global shared model to each of the clients;

[0169] The client is used to receive the global shared model sent by the server;

[0170] The client is further used to train the local model according to the global shared model and the local data to obtain a trained local model;

[0171] The client is further used to train the global shared model according to the trained local model and the local data to obtain a preliminarily trained global shared model;

[0172] The client is also used to send the global shared model that has been initially trained to the server;

[0173] The server is further configured to receive the global shared model after preliminary training sent by each of the clients;

[0174] The server is further configured to, when the set of global shared models trained in the previous round is not empty, train each of the global shared models trained in the previous round according to the global shared model, the set of global shared models trained in the previous round and the training data to obtain a set of global shared models trained;

[0175] The server side is also used to update the global shared model based on the trained global shared model set.

[0176] Reference Figure 8, shows a computer device of a model training method based on specific federated learning of the present application, which may specifically include the following:

[0177] The computer device 12 is in the form of a general-purpose computing device, and the components of the computer device 12 may include but are not limited to: one or more processors or processing units 16, a memory 28, and a bus 18 connecting different system components (including the memory 28 and the processing unit 16).

[0178] The bus 18 represents one or more of several types of bus 18 structures, including a memory bus 18 or memory controller, a peripheral bus 18, an accelerated graphics port, a processor or a local bus 18 using any of a variety of bus 18 architectures. These architectures include, by way of example, but are not limited to, an Industry Standard Architecture (ISA) bus 18, a Micro Channel Architecture (MAC) bus 18, an Enhanced ISA bus 18, an Audio Video Electronics Standards Association (VESA) local bus 18, and a Peripheral Component Interconnect (PCI) bus 18.

[0179] The computer device 12 typically includes a variety of computer system readable media. These media can be any available media that can be accessed by the computer device 12, including volatile and non-volatile media, removable and non-removable media.

[0180] Memory 28 may include computer system readable media in the form of volatile memory, such as random access memory 30 and / or cache memory 32. Computer device 12 may further include other removable / non-removable, volatile / non-volatile computer system storage media. By way of example only, storage system 34 may be used to read and write to non-removable, non-volatile magnetic media (commonly referred to as a "hard drive"). Although Figure 8 Not shown, a disk drive for reading and writing to a removable non-volatile disk (such as a "floppy disk"), and an optical disk drive for reading and writing to a removable non-volatile optical disk (such as a CD-ROM, DVD-ROM or other optical medium) may be provided. In these cases, each drive may be connected to the bus 18 via one or more data medium interfaces. The memory may include at least one program product having a set (e.g., at least one) of program modules 42, which are configured to perform the functions of the various embodiments of the present application.

[0181] A program / utility 40 having a set (at least one) of program modules 42 may be stored in, for example, a memory, such program modules 42 including, but not limited to, an operating system, one or more application programs, other program modules 42, and program data, each of which or some combination may include an implementation of a network environment. The program modules 42 generally perform the functions and / or methods of the embodiments described herein.

[0182] The computer device 12 may also communicate with one or more external devices 14 (e.g., keyboards, pointing devices, displays 24, cameras, etc.), one or more devices that enable an operator to interact with the computer device 12, and / or any device that enables the computer device 12 to communicate with one or more other computing devices (e.g., network cards, modems, etc.). Such communication may be performed through the I / O interface 22. Furthermore, the computer device 12 may also communicate with one or more networks (e.g., local area networks (LANs)), wide area networks (WANs), and / or public networks (e.g., the Internet) through a network adapter 20. Figure 8 As shown, the network adapter 20 communicates with other modules of the computer device 12 via the bus 18. It should be understood that although Figure 8 Not shown, other hardware and / or software modules may be used in conjunction with the computer device 12, including but not limited to: microcode, device drivers, redundant processing units 16, external disk drive arrays, RAID systems, tape drives, and data backup storage systems 34, etc.

[0183] The processing unit 16 executes various functional applications and data processing by running the programs stored in the memory 28, such as implementing a model training method based on specific federated learning provided in an embodiment of the present application.

[0184] That is, when the processing unit 16 executes the above program, it realizes: sending the global shared model to each of the clients; the client is used to receive the global shared model sent by the server, train the local model according to the global shared model and the local data to obtain a trained local model, train the global shared model according to the trained local model and the local data to obtain a global shared model that has been preliminarily trained, and send the preliminarily trained global shared model to the processing unit 16; receive the preliminarily trained global shared model sent by each of the clients; when the set of global shared models trained in the previous round is not empty, train each of the preliminarily trained global shared models according to the global shared model, the set of global shared models trained in the previous round and the training data to obtain a set of trained global shared models; update the global shared model according to the trained set of global shared models.

[0185] In one embodiment of the present application, a computer-readable storage medium is also provided, on which a computer program is stored. When the program is executed by a processor, a model training method based on specific federated learning as provided in all embodiments of the present application is implemented.

[0186] That is, when the program is executed by the processor, it is implemented as follows: sending the global shared model to each of the clients; the client is used to receive the global shared model sent by the server, training the local model according to the global shared model and the local data to obtain a trained local model, training the global shared model according to the trained local model and the local data to obtain a preliminarily trained global shared model, and sending the preliminarily trained global shared model to the computer-readable storage medium; receiving the preliminarily trained global shared model sent by each of the clients; when the set of global shared models trained in the previous round is not empty, training each of the preliminarily trained global shared models according to the global shared model, the set of global shared models trained in the previous round and the training data to obtain a set of trained global shared models; and updating the global shared model according to the trained set of global shared models.

[0187] Any combination of one or more computer-readable media may be used. A computer-readable medium may be a computer-readable signal medium or a computer-readable storage medium. A computer-readable storage medium may be, for example, but not limited to, an electrical, magnetic, optical, electromagnetic, infrared, or semiconductor system, device, or device, or any combination thereof. More specific examples of computer-readable storage media (a non-exhaustive list) include: an electrical connection with one or more wires, a portable computer disk, a hard disk, a random access memory (RAM), a read-only memory (ROM), an erasable programmable read-only memory (EPROM or flash memory), an optical fiber, a portable compact disk read-only memory (CD-ROM), an optical storage device, a magnetic storage device, or any suitable combination thereof. In this document, a computer-readable storage medium may be any tangible medium containing or storing a program that may be used by or in conjunction with an instruction execution system, device, or device.

[0188] Computer-readable signal media may include a data signal propagated in baseband or as part of a carrier wave, which carries a computer-readable program code. Such propagated data signals may take a variety of forms, including, but not limited to, electromagnetic signals, optical signals, or any suitable combination of the above. Computer-readable signal media may also be any computer-readable medium other than a computer-readable storage medium, which may send, propagate, or transmit a program for use by or in conjunction with an instruction execution system, apparatus, or device.

[0189] The computer program code for performing the operation of the present application can be written in one or more programming languages ​​or a combination thereof, including object-oriented programming languages, such as Java, Smalltalk, C++, and conventional procedural programming languages, such as "C" language or similar programming languages. The program code can be executed entirely on the operator's computer, partially on the operator's computer, as an independent software package, partially on the operator's computer, partially on the remote computer, or entirely on the remote computer or server. In the case of a remote computer, the remote computer can be connected to the operator's computer through any type of network, including a local area network (LAN) or a wide area network (WAN), or can be connected to an external computer (for example, using an Internet service provider to connect through the Internet). The various embodiments in this specification are described in a progressive manner, and each embodiment focuses on the differences from other embodiments. The same and similar parts between the various embodiments can be referred to each other.

[0190] Although the preferred embodiments of the present application have been described, those skilled in the art may make additional changes and modifications to these embodiments once they have learned the basic creative concept. Therefore, the appended claims are intended to be interpreted as including the preferred embodiments and all changes and modifications that fall within the scope of the embodiments of the present application.

[0191] Finally, it should be noted that, in this article, relational terms such as first and second, etc. are only used to distinguish one entity or operation from another entity or operation, and do not necessarily require or imply any such actual relationship or order between these entities or operations. Moreover, the terms "include", "comprise" or any other variants thereof are intended to cover non-exclusive inclusion, so that a process, method, article or terminal device including a series of elements includes not only those elements, but also other elements not explicitly listed, or also includes elements inherent to such process, method, article or terminal device. In the absence of further restrictions, the elements defined by the sentence "comprise a ..." do not exclude the existence of other identical elements in the process, method, article or terminal device including the elements.

[0192] The above is a detailed introduction to the model training method and device based on specific federated learning provided by the present application. This article uses specific examples to illustrate the principles and implementation methods of the present application. The description of the above embodiments is only used to help understand the method of the present application and its core idea; at the same time, for general technical personnel in this field, according to the idea of ​​the present application, there will be changes in the specific implementation method and application scope. In summary, the content of this specification should not be understood as a limitation on the present application.

Claims

1. A model training method based on specific federated learning, used in a machine learning system, wherein the machine learning system includes a server and at least two clients; The server side stores a global shared model, a global shared model set completed in a previous round of training, and training data. For the first round of training, the global shared model set completed in the previous round of training is an empty set; Each of the clients stores a local model and local data respectively; the model training method is for the server; and the model training method includes: The server sends the global shared model to each of the clients; the client is used to receive the global shared model sent by the server; the local model is trained according to the global shared model and the local data to obtain a trained local model; the global shared model is trained according to the trained local model and the local data to obtain a preliminarily trained global shared model; the preliminarily trained global shared model is sent to the server; The server receives the global shared model after preliminary training sent by each of the clients; When the global shared model set completed in the previous round of training is not empty, the server trains each of the global shared models completed in the preliminary training according to the global shared model, the global shared model set completed in the previous round of training and the training data to obtain a global shared model set completed in the training; The server updates the global shared model according to the trained global shared model set.

2. The model training method according to claim 1, characterized in that: After the server receives the global shared model sent by each client after preliminary training, the step further includes: When the set of global shared models that have completed the previous round of training is empty, the server trains each of the global shared models that have completed the preliminary training according to the training data to obtain a set of global shared models that have completed the training; The server updates the global shared model according to the trained global shared model set.

3. The model training method according to claim 1, characterized in that: The trained global shared model set includes all trained global shared models; the server trains each of the preliminarily trained global shared models according to the global shared model, the set of global shared models trained in the previous round, and the training data to obtain the trained global shared model set, including: For each global shared model that has completed preliminary training, the server performs the following steps: The server processes the training data according to the global shared model to obtain a first prediction result; The server processes the training data according to the global shared model set completed in the previous round of training to obtain a second prediction result set; The server processes the training data according to the global shared model that has been preliminarily trained to obtain a third prediction result; The server determines a first loss value according to the first prediction result, the second prediction result set, the third prediction result and a pre-constructed weighted comparison regularization loss function; The server determines a second loss value according to the third prediction result and a pre-constructed supervised reconstruction loss function; The server side trains the preliminarily trained global shared model according to the first loss value and the second loss value to obtain the trained global shared model.

4. The model training method according to claim 1, characterized in that: The trained global shared model set includes all trained global shared models; and the server updates the global shared model according to the trained global shared model set, including: The server side sets the average value of all the trained global shared models as the global shared model.

5. A model training method based on specific federated learning, used in a machine learning system, wherein the machine learning system includes a server and at least two clients; The server stores a global shared model, a global shared model set completed in the previous round of training, and training data. For the first round of training, the global shared model set completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively; the model training method is for any one of the at least two clients; it is characterized in that The model training method comprises: The client receives the global shared model sent by the server; The client trains the local model according to the global shared model and the local data to obtain a trained local model; The client trains the global shared model according to the trained local model and the local data to obtain a preliminarily trained global shared model; The client sends the global shared model that has completed preliminary training to the server; the server is used to receive the global shared model that has completed preliminary training sent by each of the clients; when the set of global shared models that have completed the previous round of training is not empty, each global shared model that has completed preliminary training is trained according to the global shared model, the set of global shared models that have completed the previous round of training and the training data to obtain a set of trained global shared models; based on the set of trained global shared models, the global shared model set is updated.

6. The model training method according to claim 5, characterized in that: The client trains the local model according to the global shared model and the local data to obtain a trained local model, including: The client processes the local data according to the global shared model to obtain a fourth prediction result; The client processes the local data according to the local model to obtain a fifth prediction result; The client determines a third loss value according to the fourth prediction result, the fifth prediction result and a pre-constructed local loss function; The client trains the local model according to the third loss value to obtain the trained local model.

7. The model training method according to claim 5, characterized in that: The step of the client training the global shared model according to the trained local model and the local data to obtain a preliminarily trained global shared model includes: The client processes the local data according to the trained local model to obtain a sixth prediction result; The client processes the local data according to the global shared model to obtain a seventh prediction result; The client determines a fourth loss value according to the sixth prediction result, the seventh prediction result and a pre-constructed shared loss function; The client trains the global shared model according to the fourth loss value to obtain the trained global shared model.

8. A model training device based on specific federated learning, used in a machine learning system, wherein the machine learning system includes a server and at least two clients; The server side stores a global shared model, a global shared model set completed in a previous round of training, and training data. For the first round of training, the global shared model set completed in the previous round of training is an empty set; Each of the clients stores a local model and local data respectively; the model training device is for the server; and the model training device comprises: A global shared model sending module, used to send the global shared model to each of the clients; the client is used to receive the global shared model sent by the server; the local model is trained according to the global shared model and the local data to obtain a trained local model; the global shared model is trained according to the trained local model and the local data to obtain a preliminarily trained global shared model; the preliminarily trained global shared model is sent to the server; A primary model receiving module, used for receiving the global shared model after preliminary training sent by each of the clients; A primary model training module, used for training each of the global shared models that have been preliminarily trained according to the global shared model, the global shared model set that has been trained in the previous round and the training data to obtain a global shared model set that has been trained, when the global shared model set that has been trained in the previous round is not empty; The global model determination module is used to update the global shared model according to the trained global shared model set.

9. A model training device based on specific federated learning, used in a machine learning system, wherein the machine learning system includes a server and at least two clients; The server stores a global shared model, a global shared model set completed in the previous round of training, and training data. For the first round of training, the global shared model set completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively; the model training device is for any one of the at least two clients; it is characterized in that The model training device comprises: A global shared model receiving module, used for receiving the global shared model sent by the server; A local model training module, used to train the local model according to the global shared model and the local data to obtain a trained local model; A global shared model training module, used to train the global shared model based on the trained local model and the local data to obtain a global shared model that has been preliminarily trained; A primary model sending module is used to send the global shared model that has completed preliminary training to the server end; the server end is used to receive the global shared model that has completed preliminary training sent by each of the clients; when the set of global shared models that have completed the previous round of training is not empty, each of the global shared models that have completed preliminary training is trained according to the global shared model, the set of global shared models that have completed the previous round of training and the training data to obtain a set of trained global shared models; based on the set of trained global shared models, the global shared model set is updated.

10. A machine learning system, characterized in that: It includes a server and at least two clients; the server stores a global shared model, a set of global shared models completed in the previous round of training, and training data. For the first round of training, the set of global shared models completed in the previous round of training is an empty set; each of the clients stores a local model and local data respectively; The server is used to send the global shared model to each of the clients; The client is used to receive the global shared model sent by the server; The client is further used to train the local model according to the global shared model and the local data to obtain a trained local model; The client is further used to train the global shared model according to the trained local model and the local data to obtain a preliminarily trained global shared model; The client is also used to send the global shared model that has been initially trained to the server; The server is further configured to receive the global shared model after preliminary training sent by each of the clients; The server is further configured to, when the set of global shared models trained in the previous round is not empty, train each of the global shared models trained in the previous round according to the global shared model, the set of global shared models trained in the previous round and the training data to obtain a set of global shared models trained; The server side is also used to update the global shared model based on the trained global shared model set.

Citation Information

Patent Citations

  • Model training method and device, electronic equipment and machine readable storage medium

    CN111967607A

  • Federation learning model training method and device and federation learning system

    CN112232528A