Model training method and system, computer equipment and storage medium
By using split neural network technology in federated learning, the local model is split into private and public parts, and the problem of high security and cost of data transmission in federated learning is solved, and efficient and secure model training is achieved.
Patent Information
- Application Number
- CN202411766303.4
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2024-12-03
- Publication Date
- 2025-05-06
AI Technical Summary
During the federated learning process, data transmission between the client and the server is vulnerable to confidentiality attacks, and the processing cost of encrypted data is high.
Through splitting neural network technology, the local model is split into private models (local encoder) and public models (local classifiers and local diffusion models), and some models are trained and optimized on the client, and the server is trained and knowledge aggregated on the global model.
It realizes that while ensuring data privacy and confidentiality, it reduces the time and communication cost of encryption and decryption, and improves the efficiency and security of model training.
Smart Images

Figure CN119940473A_ABST
Abstract
Description
Technical Field
[0001] The present disclosure relates to the field of computer technology, and in particular to a model training method, system, computer device and storage medium. Background Art
[0002] The data transmission between the client and the server during the training of federated learning is vulnerable to confidentiality attacks.
[0003] In the related art, the means of resisting confidentiality attacks is to use homomorphic encryption algorithms, differential privacy, etc. to encrypt the data transmitted between the client and the server to improve the privacy of the data transmission between the client and the server during the federated learning process. However, encrypting data consumes a lot of resources and takes a long time. As a result, the cost of protecting the security of data transmission between the client and the server during the federated learning process is high. How to reduce the cost of the security of data transmission between the client and the server during the federated learning process has become a problem that needs to be solved. Summary of the invention
[0004] In view of this, the embodiments of the present disclosure provide a model training method, system, computer device and storage medium.
[0005] In a first aspect, an embodiment of the present disclosure provides a model training method, the method comprising:
[0006] For each of the multiple clients: the server initializes the weight information of the teacher classification model corresponding to the client among the multiple teacher classification models on the server to the locally optimized weight information of the local classifier on the client received from the client, wherein the client has a local encoder as a private model, a local classifier as a public model, and a local diffusion model as a public model obtained by splitting the local model on the client; initializes the weight information of the teacher diffusion model corresponding to the client among the multiple teacher diffusion models on the server to the locally optimized weight information of the local diffusion model on the client received from the client;
[0007] On the server side, train the global diffusion model and the global classification model;
[0008] Wherein, training the global diffusion model includes: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the first Gaussian noise to obtain the first feature corresponding to the teacher diffusion model; using the global diffusion model, performing multiple forward diffusions on the first feature corresponding to each teacher diffusion model to obtain the noise-added signal corresponding to each teacher diffusion model; performing multiple reverse diffusions on the noise-added signal corresponding to each teacher diffusion model to obtain the predicted noise-added amount corresponding to each teacher diffusion model; according to the predicted noise-added amount corresponding to each teacher diffusion model and the actual noise-added amount corresponding to each teacher diffusion model, determining the loss for updating the parameters of the global diffusion model, wherein the predicted noise-added amount corresponding to the teacher diffusion model includes: the predicted noise-added amount corresponding to each forward diffusion in the multiple forward diffusions performed on the first feature corresponding to the teacher diffusion model, and the actual noise-added amount corresponding to the teacher diffusion model includes: the actual noise-added amount corresponding to each forward diffusion in the multiple forward diffusions performed on the first feature corresponding to the teacher diffusion model;
[0009] Training the global classification model includes: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the second Gaussian noise to obtain the second feature corresponding to the teacher diffusion model; using the global diffusion model, performing multiple reverse diffusions on the second Gaussian noise to obtain the second feature corresponding to the global diffusion model; for each teacher diffusion model, inputting the second feature corresponding to the teacher diffusion model into the teacher classification model corresponding to the teacher diffusion model to obtain the predicted classification result output by the teacher classification model corresponding to the teacher diffusion model; inputting the second feature corresponding to the global diffusion model into the global classification model to obtain the predicted classification result output by the global classification model; determining the loss of parameters for updating the global classification model based on the predicted classification result output by the teacher classification model corresponding to each teacher diffusion model and the predicted classification result output by the global classification model.
[0010] In a possible implementation, determining the loss of parameters for updating the global classification model according to the predicted classification results output by the teacher classification model corresponding to each teacher diffusion model and the predicted classification results output by the global classification model includes:
[0011] Determine a target prediction classification result according to the prediction classification result output by the teacher classification model corresponding to each teacher diffusion model;
[0012] According to the target prediction classification result and the prediction classification result output by the global classification model, the loss of the parameters used to update the global classification model is determined.
[0013] In a possible implementation, determining the target predicted classification result according to the predicted classification result output by the teacher classification model corresponding to each teacher diffusion model includes:
[0014] The predicted classification results output by the teacher classification models corresponding to all the teacher diffusion models are averaged and pooled to obtain the target predicted classification results.
[0015] In a possible implementation, for each client: before the server initializes the weight information of the teacher classification model corresponding to the client among the multiple teacher classification models on the server to the locally optimized weight information of the local classifier on the client received from the client, it also includes:
[0016] The client trains a local classifier on the client and a local encoder on the client, including:
[0017] Input the first local data on the client into the local encoder on the client to obtain the local features output by the local encoder on the client; input the local features into the local classifier on the client to obtain the predicted category corresponding to the first local data output by the local classifier on the client; calculate the classification loss between the predicted category corresponding to the first local data and the label of the first local data; add Gaussian noise to the local features to obtain the first feature after adding noise; use the local diffusion model on the client to perform multiple reverse diffusion on the first feature after adding noise to obtain the reconstructed feature corresponding to the local feature; calculate the loss between the local feature and the reconstructed feature corresponding to the local feature; determine the total loss corresponding to the first local data according to the classification loss and the loss between the local feature and the reconstructed feature corresponding to the local feature, and the total loss is used to update the parameters of the local classifier on the client and the parameters of the local encoder on the client;
[0018] Training the local diffusion model on the client includes: the client inputs second local data on the client into a local encoder on the client to obtain features corresponding to the second local data; using the local diffusion model on the client to perform multiple forward diffusions on the features corresponding to the second local data to obtain second features after adding noise; and using the local diffusion model on the client to perform multiple reverse diffusions on the second features after adding noise to obtain predicted noise addition amounts corresponding to each forward diffusion in the multiple forward diffusions, and calculating the loss of parameters for updating the local diffusion model on the client based on the predicted noise addition amounts corresponding to each forward diffusion in the multiple forward diffusions and the actual noise addition amounts corresponding to each forward diffusion in the multiple forward diffusions.
[0019] In a possible implementation, it also includes:
[0020] When the global diffusion model training is completed and the global classification model training is completed, the optimized weight information of the global diffusion model and the optimized weight information of the global classification model are sent to each client;
[0021] For each client, the client updates the weight information of the local diffusion model on the client to the optimized weight information of the global diffusion model; for each client, the client updates the weight information of the local classifier on the client to the optimized weight information of the global classification model.
[0022] In a second aspect, an embodiment of the present disclosure provides a model training system, which includes: a client and a server; wherein the server is used for each client: the server initializes the weight information of the teacher classification model corresponding to the client among the multiple teacher classification models on the server to the locally optimized weight information of the local classifier on the client received from the client, wherein the client has a local encoder as a private model obtained by splitting the local model on the client, a local classifier as a public model, and a local diffusion model as a public model; initializes the weight information of the teacher diffusion model corresponding to the client among the multiple teacher diffusion models on the server to the locally optimized weight information of the local classifier on the client received from the client The client receives the locally optimized weight information of the local diffusion model on the client; the server is also used to train the global diffusion model and the global classification model; wherein the training of the global diffusion model includes: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the first Gaussian noise to obtain the first feature corresponding to the teacher diffusion model; using the global diffusion model, performing multiple forward diffusions on the first feature corresponding to each teacher diffusion model to obtain the noise-added signal corresponding to each teacher diffusion model; performing multiple reverse diffusions on the noise-added signal corresponding to each teacher diffusion model to obtain the predicted noise-added amount corresponding to each teacher diffusion model; according to each The predicted noise amount corresponding to each teacher diffusion model and the actual noise amount corresponding to each teacher diffusion model are used to determine the loss of parameters for updating the global diffusion model, wherein the predicted noise amount corresponding to the teacher diffusion model includes: the predicted noise amount corresponding to each forward diffusion in multiple forward diffusions for the first feature corresponding to the teacher diffusion model, and the actual noise amount corresponding to the teacher diffusion model includes: the actual noise amount corresponding to each forward diffusion in multiple forward diffusions for the first feature corresponding to the teacher diffusion model; training the global classification model includes: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the second Gaussian noise to obtain the teacher diffusion The method comprises the following steps: first, inputting the second feature corresponding to the teacher diffusion model into the teacher classification model corresponding to the teacher diffusion model to obtain the predicted classification result output by the teacher classification model corresponding to the teacher diffusion model; second, inputting the second feature corresponding to the global diffusion model into the global classification model to obtain the predicted classification result output by the global classification model; and determining the loss of parameters for updating the global classification model according to the predicted classification result output by the teacher classification model corresponding to each teacher diffusion model and the predicted classification result output by the global classification model.
[0023] In one possible implementation, the server is also used to determine the target prediction classification result based on the prediction classification result output by the teacher classification model corresponding to each teacher diffusion model; and determine the loss of parameters used to update the global classification model based on the target prediction classification result and the prediction classification result output by the global classification model.
[0024] In a possible implementation, the server is also used to average pool the predicted classification results output by the teacher classification models corresponding to all the teacher diffusion models to obtain the target predicted classification results.
[0025] In one possible implementation, the client is also used to: for each client: before the server initializes the weight information of the teacher classification model corresponding to the client among the multiple teacher classification models on the server to the locally optimized weight information of the local classifier on the client received from the client, train the local classifier on the client and the local encoder on the client, wherein training the local classifier on the client and the local encoder on the client includes: inputting the first local data on the client into the local encoder on the client to obtain the local features output by the local encoder on the client; inputting the local features into the local classifier on the client to obtain the predicted category corresponding to the first local data output by the local classifier on the client; calculating the classification loss between the predicted category corresponding to the first local data and the label of the first local data; adding Gaussian noise to the local features to obtain the first feature after adding noise; using the local diffusion model on the client to perform multiple reverse diffusion on the first feature after adding noise to obtain the reconstructed feature corresponding to the local feature; calculating the local feature and the label. The method comprises the following steps: determining a total loss corresponding to the first local data according to the classification loss, the loss between the local feature and the reconstructed feature corresponding to the local feature, and the total loss is used to update the parameters of the local classifier on the client and the parameters of the local encoder on the client; the client is also used to train the local diffusion model on the client, wherein the training of the local diffusion model on the client comprises: the client inputs the second local data on the client into the local encoder on the client to obtain the features corresponding to the second local data; performing multiple forward diffusions on the features corresponding to the second local data using the local diffusion model on the client to obtain the second features after adding noise, and performing multiple reverse diffusions on the second features after adding noise using the local diffusion model on the client to obtain the predicted noise addition amount corresponding to each forward diffusion in the multiple forward diffusions, and calculating the loss for updating the parameters of the local diffusion model on the client according to the predicted noise addition amount corresponding to each forward diffusion in the multiple forward diffusions and the actual noise addition amount corresponding to each forward diffusion in the multiple forward diffusions.
[0026] In one possible implementation, the server is further used to send the optimized weight information of the global diffusion model and the optimized weight information of the global classification model to each client when the global diffusion model training is completed and the global classification model training is completed; the client is also used to update the weight information of the local diffusion model on the client to the optimized weight information of the global diffusion model; for each client, the client updates the weight information of the local classifier on the client to the optimized weight information of the global classification model.
[0027] In a third aspect, an embodiment of the present disclosure provides a computer device, comprising: a memory and a processor, the memory and the processor being communicatively connected to each other, computer instructions being stored in the memory, and the processor executing the method of the first aspect or any corresponding embodiment thereof by executing the computer instructions.
[0028] In a fourth aspect, an embodiment of the present disclosure provides a computer-readable storage medium having computer instructions stored thereon, the computer instructions being used to enable a computer to execute the method of the first aspect or any corresponding implementation manner thereof.
[0029] In a fifth aspect, the present invention provides a computer program product, comprising computer instructions for causing a computer to execute the method of the first aspect or any corresponding embodiment thereof.
[0030] The model training method provided by the embodiment of the present disclosure, on the one hand, is a model that takes into account both privacy and efficiency. The split neural network technology is applied to the local model on each client, and the private model in contact with the original data is retained locally, thereby maintaining the confidentiality of the model and data privacy, resisting reconstruction attacks in the modeling process, and protecting the data security of the client. Therefore, even if the parameters of the public model are directly transmitted in plain text, it is difficult for existing attack methods to infer private data only from the probability distribution, saving the time of encryption and decryption and reducing the communication cost. The split neural network technology is used to balance the privacy protection capability and model performance. Using the diffusion model as a public model, it has high operability and flexibility. Compared with other generative models, the diffusion model is easier to converge and has better operability and reproducibility. The step-by-step iterative process of adding noise and denoising in the training step controls the generated error space, allowing the model to capture real information to a greater extent, characterize richer data structures, and is more suitable for the current diverse and changing deep learning tasks.
[0031] On the other hand, the two-stage knowledge aggregation makes the global model on the server more stable and robust. Using the knowledge distillation method, the model obtained by non-IID training on the client is transferred to obtain a model with global knowledge and realize the sharing of data value. On the basis of the knowledge distillation of the response, feature distillation for the diffusion model is added to make the feature space of the global classification model input more consistent with the feature distribution of each client data. Avoid falling into the local optimum on heterogeneous data sets, so that the global model on the server can achieve a faster convergence speed.
[0032] In the first stage, the global diffusion model obtains knowledge transfer from all client encoders. In the second stage, the global classification model obtains knowledge transfer from local classifiers on all clients. By fitting the feature distribution, the global classification model can learn a common feature space from the client during the data encoding stage, ensuring that the knowledge has the same prior assumptions as each client before entering the global classification model. BRIEF DESCRIPTION OF THE DRAWINGS
[0033] In order to more clearly illustrate the specific embodiments of the present disclosure or the technical solutions in the prior art, the drawings required for use in the specific embodiments or the description of the prior art will be briefly introduced below. Obviously, the drawings described below are some embodiments of the present disclosure. For ordinary technicians in this field, other drawings can be obtained based on these drawings without paying any creative work.
[0034] Figure 1 It is a flowchart of a model training method provided by an embodiment of the present disclosure;
[0035] Figure 2 is a schematic diagram of an example of a local model on a client;
[0036] Figure 3 It is a schematic diagram of the principle of training a global diffusion model;
[0037] Figure 4 It is a schematic diagram of the principle of training a global classification model;
[0038] Figure 5 is a schematic diagram of the principle of training a local classifier on the client and a local encoder on the client;
[0039] Figure 6 It is a schematic diagram of the structure of a computer device provided in an embodiment of the present disclosure. DETAILED DESCRIPTION
[0040] In order to make the purpose, technical solution and advantages of the embodiments of the present disclosure clearer, the technical solution in the embodiments of the present disclosure will be clearly and completely described below in conjunction with the drawings in the embodiments of the present disclosure. Obviously, the described embodiments are part of the embodiments of the present disclosure, rather than all the embodiments. Based on the embodiments in the present disclosure, all other embodiments obtained by those skilled in the art without making creative work are within the scope of protection of the present disclosure.
[0041] The following first describes some concepts involved in the model training method provided in the embodiment of the present disclosure:
[0042] Horizontal Federated Learning (HFL) is a method for data collaboration and model training across multiple organizations or entities. Multiple participants in the same data feature space share models to the server and collaborate to build a global model.
[0043] Knowledge distillation is a technique to improve the performance of a student model by transferring the knowledge of a complex model (called the teacher model) to a simplified model (called the student model).
[0044] Confidentiality attack is an attack method against federated learning. The attacker attempts to obtain sensitive data or model information from the participants and only destroy the confidentiality of the model, but does not seek to destroy the federated model.
[0045] Split Learning, also known as Split Neural Network (SplitNN), is a distributed deep learning technology that divides a neural network into multiple parts, each of which is trained on a different terminal.
[0046] refer to Figure 1 , which shows an example flow chart of the model training method provided by an embodiment of the present disclosure.
[0047] In step S101, for each client: the server initializes the weight information of the teacher classification model corresponding to the client among the multiple teacher classification models on the server to the locally optimized weight information of the local classifier on the client received from the client.
[0048] In step S101, for each client: the server initializes the weight information of the teacher diffusion model corresponding to the client among the multiple teacher diffusion models on the server to the locally optimized weight information of the local diffusion model (Diffusion Model) on the client received from the client.
[0049] Among them, the client has a local encoder as a private model obtained by splitting the local model on the client, a local classifier as a public model, and a local diffusion model as a public model.
[0050] It should be noted that, for each client, the local classifier on the client, the local encoder on the client, and the local diffusion model on the client are all trained.
[0051] That is, for each client, the local classifier on the client, the local encoder on the client, and the local classifier on the client are trained before step S101.
[0052] For one client, the locally optimized weight information of the local classifier on the client is: the weight information of the local classifier on the client when the training of the local classifier on the client is completed.
[0053] For a client, the locally optimized weight information of the local diffusion model on the client is: the weight information of the local diffusion model on the client when the training of the local diffusion model on the client is completed.
[0054] In the disclosed embodiment, the multiple teacher classification models on the server correspond one-to-one to the multiple teacher diffusion models on the server.
[0055] For one client among multiple clients, the teacher classification model corresponding to the client corresponds to the teacher diffusion model corresponding to the client.
[0056] refer to Figure 2 , which is a schematic diagram showing an example of a local model on a client.
[0057] The local model on the client includes: Encoder E on the client i , classifier C on the client i , Diffusion model Df on the client i .
[0058] As an example, the encoder E on the client i In order to discard the fully connected EfficientNet-B0 network as the feature extraction backbone network, the classifier C on the client i For the MLP architecture, the diffusion model Df on the client i The denoising diffusion probabilistic model (DDPM) architecture is adopted, and the number of forward and reverse propagation times T is set to 100.
[0059] In step S102, the server trains a global diffusion model and a global classification model.
[0060] Training the global diffusion model includes: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the first Gaussian noise to obtain the first feature corresponding to the teacher diffusion model; using the global diffusion model, performing multiple forward diffusions on the first feature corresponding to each teacher diffusion model to obtain the noisy signal corresponding to each teacher diffusion model; performing multiple reverse diffusions on the noisy signal corresponding to each teacher diffusion model to obtain the predicted noise amount corresponding to each teacher diffusion model; determining the loss of the parameters used to update the global diffusion model based on the predicted noise amount corresponding to each teacher diffusion model and the actual noise amount corresponding to each teacher diffusion model.
[0061] For a teacher diffusion model, the predicted noise amount corresponding to the teacher diffusion model includes: the predicted noise amount corresponding to each forward diffusion in multiple forward diffusions corresponding to the first feature of the teacher diffusion model. The actual noise amount corresponding to the teacher diffusion model includes: the actual noise amount corresponding to each forward diffusion in multiple forward diffusions corresponding to the first feature of the teacher diffusion model. The mean mean square error (MSE) can be applied to the predicted noise amount corresponding to each forward diffusion in the multiple forward diffusions and the actual noise amount corresponding to each forward diffusion in the multiple forward diffusions to obtain the sub-loss corresponding to the teacher diffusion model for updating the parameters of the global diffusion model.
[0062] The sum of the sub-losses corresponding to each teacher diffusion model for updating the parameters of the global diffusion model may be determined as the loss for updating the parameters of the global diffusion model.
[0063] In the disclosed embodiment, when training the global diffusion model, the weight information of each teacher diffusion model is fixed to train the global diffusion model.
[0064] refer to Figure 3 , which shows a schematic diagram of the principle of training a global diffusion model.
[0065] Figure 3 The z in represents the first Gaussian noise. For Df1, Df2…Df N Each teacher diffusion model in , using the teacher diffusion model Df i , perform multiple reverse diffusions on the first Gaussian noise to obtain the first feature corresponding to the teacher diffusion model. Using the global diffusion model Df g , respectively for Df11, Df2…Df NThe first feature corresponding to each teacher diffusion model in each teacher diffusion model is forward diffused to obtain the noise-added signal corresponding to each teacher diffusion model; the noise-added signal corresponding to each teacher diffusion model is reverse diffused multiple times to obtain the predicted noise-added amount corresponding to each teacher diffusion model.
[0066] In the disclosed embodiment, training the global classification model includes: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the second Gaussian noise to obtain a second feature corresponding to the teacher diffusion model; using the global diffusion model, performing multiple reverse diffusions on the second Gaussian noise to obtain a second feature corresponding to the global diffusion model; for each teacher diffusion model, inputting the second feature corresponding to the teacher diffusion model into the teacher classification model corresponding to the teacher diffusion model to obtain a predicted classification result output by the teacher classification model corresponding to the teacher diffusion model; inputting the second feature corresponding to the global diffusion model into the global classification model to obtain a predicted classification result output by the global classification model; determining the loss of parameters for updating the global classification model based on the predicted classification result output by the teacher classification model corresponding to each teacher diffusion model and the predicted classification result output by the global classification model.
[0067] In one possible implementation, determining the loss of parameters for updating the global classification model based on the predicted classification results output by the teacher classification model corresponding to each teacher diffusion model and the predicted classification results output by the global classification model includes: determining the target predicted classification results based on the predicted classification results output by the teacher classification model corresponding to each teacher diffusion model; determining the loss of parameters for updating the global classification model based on the target predicted classification results and the predicted classification results output by the global classification model.
[0068] In one possible implementation, determining the target prediction classification result based on the prediction classification result output by the teacher classification model corresponding to each teacher diffusion model includes: averaging the prediction classification results output by the teacher classification models corresponding to all teacher diffusion models to obtain the target prediction classification result.
[0069] The predicted classification results output by the teacher classification model corresponding to all teacher diffusion models are averaged and pooled, and the target predicted classification result P can be expressed as:
[0070]
[0071] Among them, C i (E i (z) represents the predicted classification result output by the teacher classification model corresponding to the teacher diffusion model, and T is the temperature of knowledge distillation. As an example, is 7.
[0072] The predicted distribution Q corresponding to the predicted classification result output by the global classification model can be expressed as:
[0073]
[0074] Among them, E g (z) is the predicted classification result output by the global classification model.
[0075] When determining the loss for updating the parameters of the global classification model based on the target prediction classification result and the prediction classification result output by the global classification model, the KL divergence is used as the loss function of knowledge distillation to calculate the loss between the prediction distribution and the target distribution σ(P), which is the loss L for updating the parameters of the global classification model. KD , where σ is the softnax function.
[0076] The KL divergence is used to calculate the loss L between the predicted distribution and the target distribution σ(P) for updating the parameters of the global classification model. KD It can be expressed as:
[0077] L KD =KL(σ(P),,log(σ(Q)))
[0078] refer to Figure 4 , which shows a schematic diagram of the principle of training a global classification model.
[0079] For Df1, Df2…Df N For each teacher diffusion model in the , use the teacher diffusion model to perform multiple reverse diffusions on the second Gaussian noise to obtain the second feature corresponding to the teacher diffusion model; use the global diffusion model to perform multiple reverse diffusions on the second Gaussian noise to obtain the second feature corresponding to the global diffusion model; for each teacher diffusion model, input the second feature corresponding to the teacher diffusion model into the teacher classification model corresponding to the teacher diffusion model to obtain the predicted classification result output by the teacher classification model corresponding to the teacher diffusion model.
[0080] Df1, Df2…Df N The predicted classification results output by the corresponding teacher classification models are: C1, C2…C N , the global diffusion model Df g The corresponding second feature is input into the global classification model Df g , get the global classification model Df g Output prediction classification result C g ; The predicted classification results output by the teacher classification model corresponding to all teacher diffusion models, namely C1, C2…C NPerform average pooling to obtain the target prediction classification result. Using KL divergence, according to the target prediction classification result and the global classification model Df g Output prediction classification result C g , determines the loss used to update the parameters of the global classification model.
[0081] In a possible implementation, the method further includes: step S104.
[0082] In step S104, when the global diffusion model training is completed and the global classification model training is completed, the optimized weight information of the global diffusion model and the optimized weight information of the global classification model are sent to each client; for each client among the multiple clients, the client updates the weight information of the local diffusion model on the client to the optimized weight information of the global diffusion model; the client updates the weight information of the local classifier on the client to the optimized weight information of the global classification model.
[0083] The optimized weight information of the global diffusion model is: the weight information of the global diffusion model when the global diffusion model training is completed.
[0084] The optimized weight information of the global classification model is: the weight information of the global classification model when the training of the global classification model is completed.
[0085] In a possible implementation, step S100 is also included.
[0086] In step S100 , for each client, the client trains a local classifier on the client, a local encoder on the client, and a local diffusion model on the client.
[0087] For a client, training a local classifier on the client, a local encoder on the client, and a local diffusion model on the client includes step S1001 and step S1002.
[0088] refer to Figure 5 , which shows a schematic diagram of the principle of training a local classifier on the client and a local encoder on the client.
[0089] In step S1001, the client trains a local classifier on the client and a local encoder on the client. In step S1002, the client trains a local diffusion model on the client.
[0090] Step S1001 includes: step S10011-step S10013.
[0091] In step S10011, for a client, the first local data on the client is input into the local encoder on the client to obtain the local feature E output by the local encoder on the client. i (x); the local feature E output by the local encoder on the client i (x) Input to the local classifier C on the client i , get the local classifier C on the client i The output corresponds to the predicted category of the first local data; calculate the local classifier C on the client i The output corresponds to the classification loss between the predicted category of the first local and the label y of the first local data.
[0092] The cross entropy loss function can be applied to the local classifier C on the client i The output corresponding to the predicted category of the first local data and the label y of the first local data obtains the local classifier C on the client i The classification loss between the predicted category of the first local data and the label y of the first local data outputted by the local classifier C on the client i The classification loss L between the predicted category of the first local data and the label y of the first local data output cls It can be expressed as:
[0093] L cLs =CE(y,C i (E i (X)))
[0094] In step S10012, the local feature E output by the local encoder on the client is i (x) Add Gaussian noise to obtain the first feature after adding noise; use the local diffusion model Df i , for the first feature E after adding noise i (x) Perform multiple reverse diffusions, i.e., T times, to obtain the local feature E output by the local encoder on the client. i (x) The corresponding reconstructed feature; calculate the local feature E i (x) and the local feature E i (x) The loss between the corresponding reconstructed features.
[0095] Here, Gaussian noise can be represented by z.
[0096] After adding noise to the first feature E i (x) multiple times, that is, T times of reverse diffusion, can be expressed as:
[0097]
[0098] When calculating the local feature E i (x) and local feature E i When the loss between the corresponding reconstruction features (x) is calculated, the mean square error (MSE) function can be applied to the local feature E i (x) and local feature E i (x) The corresponding reconstructed feature x0 is obtained to obtain the local feature E i (x) and local feature E i (x) The loss between the corresponding reconstructed features x0.
[0099] Apply the mean square error (MSE) function to the local features E i (x) and local feature E i (x) The corresponding reconstructed feature x0 is obtained to obtain the local feature E i (x) and local feature E i (x) The loss L between the corresponding reconstructed features x0 rec It can be expressed as:
[0100] L rec =MSE(E i (x),x0)
[0101] In step S10013, according to the classification loss, the local feature E i (x) and the local feature E t (x) The loss between the corresponding reconstructed features is used to determine the total loss corresponding to the first local data, and the total loss is used to update the parameters of the local classifier on the client and the parameters of the local encoder on the client.
[0102] The total loss L can be expressed as:
[0103] L=L cls +αL rec
[0104] Among them, α is a hyperparameter.
[0105] In step S1002, the client inputs the second local data on the client into the local encoder on the client to obtain features corresponding to the second local data; uses the local diffusion model on the client to perform multiple forward diffusions, i.e., T times, on the features corresponding to the second local data to obtain second features after adding noise; and uses the local diffusion model on the client to perform T times of reverse diffusion on the second features after adding noise to obtain a predicted amount of noise added corresponding to each forward diffusion in the T times of forward diffusion, and calculates the loss of parameters for updating the local diffusion model on the client based on the predicted amount of noise added corresponding to each forward diffusion in the T times of forward diffusion and the actual amount of noise added corresponding to each forward diffusion in the T times of forward diffusion.
[0106] It should be noted that the first local data and the second local data may be the same.
[0107] In the embodiment of the present disclosure, when training the local diffusion model on the client, the weight of the local encoder on the client is fixed.
[0108] The local diffusion model on the client performs forward diffusion multiple times, that is, T times, on the features corresponding to the second local data to obtain the second feature x after adding noise. t It can be expressed as:
[0109]
[0110] Among them, x0 is the feature corresponding to the second local data, ∈ conforms to the Gaussian distribution, Indicates the cumulative multiplication of the noise level parameters.
[0111] When calculating the loss for updating the parameters of the local diffusion model on the client according to the predicted noise addition amount corresponding to each forward diffusion in the T forward diffusions and the actual noise addition amount corresponding to each forward diffusion in the T forward diffusions, the mean mean square error (MSE) can be applied to the predicted noise addition amount corresponding to each forward diffusion in the T forward diffusions and the actual noise addition amount corresponding to each forward diffusion in the T forward diffusions to obtain the loss for updating the parameters of the local diffusion model on the client.
[0112] The loss L used to update the parameters of the local diffusion model t It can be expressed as:
[0113] L t =MSE(∈-∈ θ (x t , t))
[0114] Among them, ∈ is the actual noise amount corresponding to the t-th step of forward diffusion in the T-th forward diffusion, ∈ θ (x t ,t) is the predicted noise amount corresponding to the t-th forward diffusion in the T forward diffusions predicted by the U-Net network.
[0115] The embodiments of the present disclosure provide a model training system. The system is used to implement the above embodiments and preferred implementation modes, and the descriptions that have been made will not be repeated. As used below, the term "unit" can implement a combination of software and / or hardware for a predetermined function. Although the devices described in the following embodiments are preferably implemented in software, the implementation of hardware, or a combination of software and hardware, is also possible and conceivable.
[0116] The model training system includes: a client and a server; wherein the server is used for each client: the server initializes the weight information of the teacher classification model corresponding to the client among multiple teacher classification models on the server to the locally optimized weight information of the local classifier on the client received from the client, wherein the client has a local encoder as a private model obtained by splitting the local model on the client, a local classifier as a public model, and a local diffusion model as a public model; initializes the weight information of the teacher diffusion model corresponding to the client among multiple teacher diffusion models on the server to the local diffusion model on the client received from the client The server is also used to train the global diffusion model and the global classification model; wherein the training of the global diffusion model comprises: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the first Gaussian noise to obtain the first feature corresponding to the teacher diffusion model; using the global diffusion model, performing multiple forward diffusions on the first feature corresponding to each teacher diffusion model to obtain the noise-added signal corresponding to each teacher diffusion model; performing multiple reverse diffusions on the noise-added signal corresponding to each teacher diffusion model to obtain the predicted noise-added amount corresponding to each teacher diffusion model; and performing multiple reverse diffusions on the predicted noise-added amount corresponding to each teacher diffusion model according to the predicted noise-added amount corresponding to each teacher diffusion model. The noise amount and the actual noise amount corresponding to each teacher diffusion model are measured to determine the loss of parameters for updating the global diffusion model, wherein the predicted noise amount corresponding to the teacher diffusion model includes: the predicted noise amount corresponding to each forward diffusion in multiple forward diffusions for the first feature corresponding to the teacher diffusion model, and the actual noise amount corresponding to the teacher diffusion model includes: the actual noise amount corresponding to each forward diffusion in multiple forward diffusions for the first feature corresponding to the teacher diffusion model; training the global classification model includes: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the second Gaussian noise to obtain the first feature corresponding to the teacher diffusion model Second feature; using the global diffusion model, the second Gaussian noise is reversely diffused multiple times to obtain the second feature corresponding to the global diffusion model; for each teacher diffusion model, the second feature corresponding to the teacher diffusion model is input into the teacher classification model corresponding to the teacher diffusion model to obtain the predicted classification result output by the teacher classification model corresponding to the teacher diffusion model; the second feature corresponding to the global diffusion model is input into the global classification model to obtain the predicted classification result output by the global classification model; according to the predicted classification result output by the teacher classification model corresponding to each teacher diffusion model and the predicted classification result output by the global classification model, the loss of the parameters used to update the global classification model is determined.
[0117] In one possible implementation, the server is also used to determine the target prediction classification result based on the prediction classification result output by the teacher classification model corresponding to each teacher diffusion model; and determine the loss of parameters used to update the global classification model based on the target prediction classification result and the prediction classification result output by the global classification model.
[0118] In a possible implementation, the server is also used to average pool the predicted classification results output by the teacher classification models corresponding to all the teacher diffusion models to obtain the target predicted classification results.
[0119] In one possible implementation, the client is also used to: for each client: before the server initializes the weight information of the teacher classification model corresponding to the client among the multiple teacher classification models on the server to the locally optimized weight information of the local classifier on the client received from the client, train the local classifier on the client and the local encoder on the client, wherein training the local classifier on the client and the local encoder on the client includes: inputting the first local data on the client into the local encoder on the client to obtain the local features output by the local encoder on the client; inputting the local features into the local classifier on the client to obtain the predicted category corresponding to the first local data output by the local classifier on the client; calculating the classification loss between the predicted category corresponding to the first local data and the label of the first local data; adding Gaussian noise to the local features to obtain the first feature after adding noise; using the local diffusion model on the client to perform multiple reverse diffusion on the first feature after adding noise to obtain the reconstructed feature corresponding to the local feature; calculating the local feature and the label. The method comprises the following steps: determining a total loss corresponding to the first local data according to the classification loss, the loss between the local feature and the reconstructed feature corresponding to the local feature, and the total loss is used to update the parameters of the local classifier on the client and the parameters of the local encoder on the client; the client is also used to train the local diffusion model on the client, wherein the training of the local diffusion model on the client comprises: the client inputs the second local data on the client into the local encoder on the client to obtain the features corresponding to the second local data; performing multiple forward diffusions on the features corresponding to the second local data using the local diffusion model on the client to obtain the second features after adding noise, and performing multiple reverse diffusions on the second features after adding noise using the local diffusion model on the client to obtain the predicted noise addition amount corresponding to each forward diffusion in the multiple forward diffusions, and calculating the loss for updating the parameters of the local diffusion model on the client according to the predicted noise addition amount corresponding to each forward diffusion in the multiple forward diffusions and the actual noise addition amount corresponding to each forward diffusion in the multiple forward diffusions.
[0120] In one possible implementation, the server is further used to send the optimized weight information of the global diffusion model and the optimized weight information of the global classification model to each client when the global diffusion model training is completed and the global classification model training is completed; the client is also used to update the weight information of the local diffusion model on the client to the optimized weight information of the global diffusion model; for each client, the client updates the weight information of the local classifier on the client to the optimized weight information of the global classification model.
[0121] In this embodiment, the system is presented in the form of functional units, where the units refer to ASIC circuits, processors and memories that execute one or more software or fixed programs, and / or other devices that can provide the above functions.
[0122] The further functional description of each of the above units is the same as that of the above corresponding embodiments and will not be repeated here.
[0123] refer to Figure 6 , which shows a schematic diagram of the structure of a computer device provided by an embodiment of the present disclosure, the computer device includes: one or more processors 10, a memory 20, and interfaces for connecting various components, including high-speed interfaces and low-speed interfaces. The various components are connected to each other using different buses for communication, and can be installed on a common motherboard or installed in other ways as needed. The processor can process instructions executed in the computer device, including instructions stored in or on the memory to display graphical information of the GUI on an external input / output device (such as a display device coupled to the interface). In some optional embodiments, if necessary, multiple processors and / or multiple buses can be used together with multiple memories and multiple memories. Similarly, multiple computer devices can be connected, and each device provides part of the necessary operations (for example, as a server array, a group of blade servers, or a multi-processor system).
[0124] The processor 10 may be a central processing unit, a network processor or a combination thereof. The processor 10 may further include a hardware chip. The hardware chip may be a dedicated integrated circuit, a programmable logic device or a combination thereof. The programmable logic device may be a complex programmable logic device, a field programmable gate array, a general purpose array logic or any combination thereof.
[0125] The memory 20 stores instructions executable by at least one processor 10, so that the at least one processor 10 executes the method shown in the above embodiment.
[0126] The memory 20 may include a program storage area and a data storage area, wherein the program storage area may store an operating system, an application required for at least one function; the data storage area may store data created according to the use of the computer device, etc. In addition, the memory 20 may include a high-speed random access memory, and may also include a non-transient memory, such as at least one disk storage device, a flash memory device, or other non-transient solid-state storage device. In some optional embodiments, the memory 20 may optionally include a memory remotely arranged relative to the processor 10, and these remote memories may be connected to the computer device via a network. Examples of the above-mentioned network include, but are not limited to, the Internet, an intranet, a local area network, a mobile communication network, and combinations thereof.
[0127] The memory 20 may include a volatile memory, such as a random access memory; the memory may also include a non-volatile memory, such as a flash memory, a hard disk or a solid state drive; the memory 20 may also include a combination of the above types of memory.
[0128] The computer device further includes an input device 30 and an output device 40. The processor 10, the memory 20, the input device 30 and the output device 40 may be connected via a bus or other means.
[0129] The input device 30 can receive input digital or character information, and generate key signal input related to the user settings and function control of the computer device, such as a touch screen, a keypad, a mouse, a track pad, a touch pad, an indicator bar, one or more mouse buttons, a trackball, a joystick, etc. The output device 40 may include a display device, an auxiliary lighting device (e.g., an LED) and a tactile feedback device (e.g., a vibration motor), etc. The above-mentioned display device includes but is not limited to a liquid crystal display, a light emitting diode, a display and a plasma display. In some optional embodiments, the display device can be a touch screen.
[0130] The embodiments of the present disclosure also provide a computer-readable storage medium. The above-mentioned method according to the embodiments of the present disclosure can be implemented in hardware, firmware, or can be implemented as a computer code that can be recorded in a storage medium, or can be implemented as a computer code that is originally stored in a remote storage medium or a non-temporary machine-readable storage medium and will be stored in a local storage medium and downloaded through a network, so that the method described herein can be stored in such software processing on a storage medium using a general-purpose computer, a dedicated processor, or programmable or dedicated hardware. Among them, the storage medium can be a magnetic disk, an optical disk, a read-only storage memory, a random access memory, a flash memory, a hard disk or a solid-state drive, etc.; further, the storage medium can also include a combination of the above-mentioned types of memory. It can be understood that a computer, a processor, a microprocessor controller, or programmable hardware includes a storage component that can store or receive software or computer code. When the software or computer code is accessed and executed by a computer, a processor, or hardware, the method shown in the above embodiment is implemented.
[0131] A portion of the embodiments of the present disclosure may be applied as a computer program product, such as a computer program instruction, which, when executed by a computer, can call or provide the method and / or technical solution according to the present invention through the operation of the computer. Those skilled in the art should understand that the existence of computer program instructions in computer-readable media includes, but is not limited to, source files, executable files, installation package files, etc., and accordingly, the way in which computer program instructions are executed by a computer includes, but is not limited to: the computer directly executes the instruction, or the computer compiles the instruction and then executes the corresponding compiled program, or the computer reads and executes the instruction, or the computer reads and installs the instruction and then executes the corresponding installed program. Here, the computer-readable medium may be any available computer-readable storage medium or communication medium accessible to the computer.
[0132] Although the embodiments of the present disclosure have been described in conjunction with the accompanying drawings, those skilled in the art may make various modifications and variations without departing from the spirit and scope of the present disclosure, and such modifications and variations are all within the scope defined by the appended claims.
Claims
1. A model training method, characterized in that: The method comprises: For each of the multiple clients: the server initializes the weight information of the teacher classification model corresponding to the client among the multiple teacher classification models on the server to the locally optimized weight information of the local classifier on the client received from the client, wherein the client has a local encoder as a private model, a local classifier as a public model, and a local diffusion model as a public model obtained by splitting the local model on the client; initializes the weight information of the teacher diffusion model corresponding to the client among the multiple teacher diffusion models on the server to the locally optimized weight information of the local diffusion model on the client received from the client; On the server side, train the global diffusion model and the global classification model; Wherein, training the global diffusion model includes: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the first Gaussian noise to obtain the first feature corresponding to the teacher diffusion model; using the global diffusion model, performing multiple forward diffusions on the first feature corresponding to each teacher diffusion model to obtain the noise-added signal corresponding to each teacher diffusion model; performing multiple reverse diffusions on the noise-added signal corresponding to each teacher diffusion model to obtain the predicted noise-added amount corresponding to each teacher diffusion model; according to the predicted noise-added amount corresponding to each teacher diffusion model and the actual noise-added amount corresponding to each teacher diffusion model, determining the loss for updating the parameters of the global diffusion model, wherein the predicted noise-added amount corresponding to the teacher diffusion model includes: the predicted noise-added amount corresponding to each forward diffusion in the multiple forward diffusions performed on the first feature corresponding to the teacher diffusion model, and the actual noise-added amount corresponding to the teacher diffusion model includes: the actual noise-added amount corresponding to each forward diffusion in the multiple forward diffusions performed on the first feature corresponding to the teacher diffusion model; Training the global classification model includes: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the second Gaussian noise to obtain the second feature corresponding to the teacher diffusion model; using the global diffusion model, performing multiple reverse diffusions on the second Gaussian noise to obtain the second feature corresponding to the global diffusion model; for each teacher diffusion model, inputting the second feature corresponding to the teacher diffusion model into the teacher classification model corresponding to the teacher diffusion model to obtain the predicted classification result output by the teacher classification model corresponding to the teacher diffusion model; inputting the second feature corresponding to the global diffusion model into the global classification model to obtain the predicted classification result output by the global classification model; determining the loss of parameters for updating the global classification model based on the predicted classification result output by the teacher classification model corresponding to each teacher diffusion model and the predicted classification result output by the global classification model.
2. The method according to claim 1, characterized in that Determining the loss of parameters for updating the global classification model according to the predicted classification results output by the teacher classification model corresponding to each teacher diffusion model and the predicted classification results output by the global classification model includes: Determine a target prediction classification result according to the prediction classification result output by the teacher classification model corresponding to each teacher diffusion model; According to the target prediction classification result and the prediction classification result output by the global classification model, the loss of the parameters used to update the global classification model is determined.
3. The method according to claim 2, characterized in that According to the predicted classification results output by the teacher classification model corresponding to each teacher diffusion model, determining the target predicted classification results includes: The predicted classification results output by the teacher classification models corresponding to all the teacher diffusion models are averaged and pooled to obtain the target predicted classification results.
4. The method according to claim 1, characterized in that: Before the server initializes, for each of the multiple clients, the weight information of the teacher classification model corresponding to the client among the multiple teacher classification models on the server to the locally optimized weight information of the local classifier on the client received from the client, the method further comprises: The client trains a local classifier on the client and a local encoder on the client, including: Input the first local data on the client into the local encoder on the client to obtain the local features output by the local encoder on the client; input the local features into the local classifier on the client to obtain the predicted category corresponding to the first local data output by the local classifier on the client; calculate the classification loss between the predicted category corresponding to the first local data and the label of the first local data; add Gaussian noise to the local features to obtain the first feature after adding noise; use the local diffusion model on the client to perform multiple reverse diffusion on the first feature after adding noise to obtain the reconstructed feature corresponding to the local feature; calculate the loss between the local feature and the reconstructed feature corresponding to the local feature; determine the total loss corresponding to the first local data according to the classification loss and the loss between the local feature and the reconstructed feature corresponding to the local feature, and the total loss is used to update the parameters of the local classifier on the client and the parameters of the local encoder on the client; The client trains the local diffusion model on the client, including: the client inputs second local data on the client into a local encoder on the client to obtain features corresponding to the second local data; uses the local diffusion model on the client to perform multiple forward diffusions on the features corresponding to the second local data to obtain second features after adding noise; and uses the local diffusion model on the client to perform multiple reverse diffusions on the second features after adding noise to obtain predicted noise addition amounts corresponding to each forward diffusion in the multiple forward diffusions, and calculates the loss of parameters for updating the local diffusion model on the client according to the predicted noise addition amounts corresponding to each forward diffusion in the multiple forward diffusions and the actual noise addition amounts corresponding to each forward diffusion in the multiple forward diffusions.
5. The method according to any one of claims 1 to 4, characterized in that The method further comprises: When the global diffusion model training is completed and the global classification model training is completed, the server sends the optimized weight information of the global diffusion model and the optimized weight information of the global classification model to each of the clients; For each client, the client updates the weight information of the local diffusion model on the client to the optimized weight information of the global diffusion model; For each client, the client updates the weight information of the local classifier on the client to the optimized weight information of the global classification model.
6. A model training system, characterized in that: The system includes: a client and a server; wherein the server is used for each client: the server initializes the weight information of the teacher classification model corresponding to the client among multiple teacher classification models on the server to the locally optimized weight information of the local classifier on the client received from the client, wherein the client has a local encoder as a private model obtained by splitting the local model on the client, a local classifier as a public model, and a local diffusion model as a public model; initializes the weight information of the teacher diffusion model corresponding to the client among multiple teacher diffusion models on the server to the local diffusion model on the client received from the client The server is also used to train the global diffusion model and the global classification model; wherein the training of the global diffusion model comprises: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the first Gaussian noise to obtain the first feature corresponding to the teacher diffusion model; using the global diffusion model, performing multiple forward diffusions on the first feature corresponding to each teacher diffusion model to obtain the noise-added signal corresponding to each teacher diffusion model; performing multiple reverse diffusions on the noise-added signal corresponding to each teacher diffusion model to obtain the predicted noise-added amount corresponding to each teacher diffusion model; according to the predicted value corresponding to each teacher diffusion model The noise amount and the actual noise amount corresponding to each teacher diffusion model are used to determine the loss of parameters for updating the global diffusion model, wherein the predicted noise amount corresponding to the teacher diffusion model includes: the predicted noise amount corresponding to each forward diffusion in multiple forward diffusions of the first feature corresponding to the teacher diffusion model, and the actual noise amount corresponding to the teacher diffusion model includes: the actual noise amount corresponding to each forward diffusion in multiple forward diffusions of the first feature corresponding to the teacher diffusion model; training the global classification model includes: for each teacher diffusion model, using the teacher diffusion model, performing multiple reverse diffusions on the second Gaussian noise to obtain the first feature corresponding to the teacher diffusion model Second feature; using the global diffusion model, the second Gaussian noise is reversely diffused multiple times to obtain the second feature corresponding to the global diffusion model; for each teacher diffusion model, the second feature corresponding to the teacher diffusion model is input into the teacher classification model corresponding to the teacher diffusion model to obtain the predicted classification result output by the teacher classification model corresponding to the teacher diffusion model; the second feature corresponding to the global diffusion model is input into the global classification model to obtain the predicted classification result output by the global classification model; according to the predicted classification result output by the teacher classification model corresponding to each teacher diffusion model and the predicted classification result output by the global classification model, the loss of the parameters used to update the global classification model is determined.
7. The system according to claim 6, characterized in that The server is also used to determine the target prediction classification result based on the prediction classification result output by the teacher classification model corresponding to each teacher diffusion model; and determine the loss of parameters used to update the global classification model based on the target prediction classification result and the prediction classification result output by the global classification model.
8. A computer device, characterized in that: include: A memory and a processor, wherein the memory and the processor are communicatively connected to each other, the memory stores computer instructions, and the processor executes the method according to any one of claims 1 to 5 by executing the computer instructions.
9. A computer-readable storage medium, characterized in that: The computer-readable storage medium stores computer instructions, and the computer instructions are used to enable a computer to execute the method according to any one of claims 1 to 5.
10. A computer program product, characterized in that The method comprises computer instructions for causing a computer to execute the method according to any one of claims 1 to 5.