A federated learning method based on Wiener deconvolution and diagonal Fisher pruning
By adopting Diag-Fisher pruning and natural gradient optimization methods in federated learning, combined with Wiener deconvolution denoising processing, the problems of noise accumulation and communication overhead caused by homomorphic encryption are solved, and more efficient model training and better accuracy are achieved.
Patent Information
- Application Number
- CN202411696958.9
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2024-11-26
- Publication Date
- 2025-05-06
- Estimated Expiration
- 2044-11-26
AI Technical Summary
In the existing federated learning technology, homomorphic encryption leads to noise accumulation affecting model accuracy, and frequent parameters and gradient exchanges lead to huge communication overhead.
The Diag-Fisher pruning mechanism is used to prune the gradient through the Fisher information matrix, and combined with the natural gradient optimization method, the client uploads the pruned gradient to the server through homomorphic encryption. The server side uses a federated mechanism based on Wiener deconvolution to denoise the encryption gradient, reducing noise accumulation and enhancing useful signals.
It effectively reduces communication overhead, reduces the impact of noise on model accuracy, and improves the convergence speed and final training accuracy of the model.
Smart Images

Figure CN119180353B_ABST
Abstract
Description
Technical Field
[0001] The invention relates to a federated learning method based on Wiener deconvolution and diagonal Fisher pruning, belonging to the technical field of federated learning. Background Art
[0002] In recent years, with the rapid development of artificial intelligence technology, data from various fields such as finance, manufacturing, and services have been widely used in artificial intelligence model training. Traditional artificial intelligence models need to centralize data on the server side for learning and training, which can easily lead to user privacy leakage. In order to solve this problem, Google proposed the concept of Federated Learning (FL) in 2016. Data holders train models locally and then upload model training parameters to the server for aggregate updates, thereby avoiding the leakage of private data. However, relevant studies have shown that it is still possible to restore some of the original data of users or determine the relevant attribute characteristics of specific users by analyzing the intermediate parameters uploaded by users. Therefore, even if the federated learning mechanism is adopted, corresponding privacy protection technology is still needed to further prevent the leakage of user privacy.
[0003] At present, the privacy protection technologies of federated learning mainly include homomorphic encryption, multi-party secure computing, and differential privacy. Compared with other technologies, homomorphic encryption can provide more reliable privacy protection services. Most of the current federated learning models use the second-generation homomorphic encryption scheme, which has greatly improved performance, efficiency, and practicality compared with the first-generation scheme, and supports the calculation of arbitrary Boolean circuits. However, although the second-generation homomorphic encryption scheme introduces module exchange technology to control the expansion of ciphertext, it cannot completely eliminate the impact of noise accumulation on model accuracy, and does not use the statistical characteristics of the signal to enhance the signal.
[0004] In addition, during the federated learning process, the user end and the server need to frequently exchange parameters and gradients, resulting in huge communication overhead. At present, strategies to reduce the communication overhead of federated learning can be divided into two types: model compression and distillation algorithms. Model compression strategies reduce the accuracy of the model and cannot handle heterogeneous data; although distillation algorithms can maintain good performance when data is heterogeneous, their communication overhead has not been effectively reduced. In recent years, researchers have proposed using pruning mechanisms to reduce the communication overhead of federated learning models, such as Federated Dropout (FedDropout). However, federated learning models using pruning mechanisms require a large amount of computation on the user end, which increases computational overhead while reducing communication overhead. Summary of the invention
[0005] The purpose of the present invention is to overcome the shortcomings of the prior art and provide a federated learning method based on Wiener deconvolution and diagonal Fisher pruning, which can reduce the noise accumulation caused by homomorphic encryption in federated learning while retaining and enhancing useful signals, and reduce communication overhead by reducing gradient upload.
[0006] After the client completes model training, it uses the Diag-Fisher pruning mechanism to prune the gradient through the Fisher information matrix. Then, the natural gradient optimization is applied to correct the pruned gradient and upload it to the server with homomorphic encryption. After receiving the encrypted pruned gradient, the server uses a federated mechanism based on Wiener deconvolution to filter the noise. By estimating the noise power spectral density and designing deconvolution, the noise introduced in the encryption process is removed and the useful signal is enhanced. Finally, the server aggregates the gradients uploaded by each client, forms a global model and feeds it back to the client to complete the training iteration of federated learning.
[0007] A federated learning method based on Wiener deconvolution and diagonal Fisher pruning, comprising the following steps:
[0008] Step 1: The client receives the global model from the server and trains it based on the local dataset.
[0009] Step 2: The client performs local training on the received global model based on the local data set to obtain a local model.
[0010] Step 3: Use the Diag-Fisher pruning mechanism to prune the gradient of the local model, calculate and retain the gradients with large information contribution in the diagonal elements of the Fisher information matrix; use the Diag-Fisher pruning mechanism to calculate the importance of the gradient through the Fisher information matrix, and retain important gradients based on the pruning threshold to reduce communication overhead. The process is as follows:
[0011] 3.1 Initialize the diagonal Fisher information matrix corresponding to the local model, and its calculation expression is: ,
[0012] 3.2 For each batch of data , forward propagation calculates the loss ,
[0013] in, For Clients The local model gradient of A batch of data,
[0014] 3.3 For each batch of data , calculate the gradient based on the loss back propagation ,
[0015] in, For Clients The local model gradient of For its data set A batch of data, represents the loss calculated by forward propagation,
[0016] 3.4 For each batch of data , the diagonal Fisher information matrix of the local model is accumulated ,in, The gradient calculated by back propagation,
[0017] 3.5 According to the batch dataset that represents the current client for updating the model The size of the normalized diagonal Fisher information matrix ,in, Represents the batch dataset used by the current client to update the model The size of
[0018] 3.6 Calculate the pruning threshold based on the set pruning rate and the calculated diagonal Fisher information matrix ,
[0019] in, is the calculated diagonal Fisher information matrix, is the pre-set pruning rate, The pruning threshold is obtained by calculation.
[0020] 3.7 Compare the Fisher information value in the diagonal Fisher information matrix with the pruning threshold to generate a mask
[0021] ,
[0022] 3.8 Apply mask to prune local model parameters: ,in, The pruning threshold is obtained by calculation. The mask is calculated, and finally the mask is applied to get the new gradient after pruning ,
[0023] Step 4: Use natural gradient optimization on the pruned local model. The process is as follows:
[0024] 4.1 For each batch of data , forward propagation calculates the loss ,
[0025] 4.2 For each batch of data , calculate the gradient based on the loss back propagation
[0026] ,
[0027] 4.3 Obtaining the natural gradient based on the diagonal Fisher information matrix and gradient calculation ,
[0028] in, For Clients The pruned local model gradients, is a small positive value, which is a smoothing term used to prevent the denominator from being zero and ensure numerical stability. is the calculated natural gradient,
[0029] 4.4 Update local model parameters according to the calculated natural gradient: ,
[0030] in, is the learning rate, In order to use the local model gradient after natural gradient update, the client uses natural gradient optimization on the pruned local model to optimize and correct the pruned gradient to ensure that the gradient update direction conforms to the geometric structure of the parameter space, thereby improving the convergence speed and accuracy of the model.
[0031] Step 5: The client uploads the pruned model parameters and optimized gradients to the server through homomorphic encryption.
[0032] Step 6: Use the federated mechanism based on Wiener deconvolution to denoise the encrypted gradients. After receiving the encrypted gradients from multiple clients, the server uses the federated mechanism based on Wiener deconvolution to denoise the encrypted gradients, estimate the noise power spectrum density, and eliminate the noise introduced in the homomorphic encryption process by designing Wiener deconvolution to enhance the effective signal. The server aggregates the denoised gradients from multiple clients to obtain the aggregated global model. The server sends the updated global model to the client to complete the iteration of the current training round. When the preset iteration round is reached, the server finally generates a global model trained by federated learning. The process is as follows:
[0033] 6.1 Fourier transform of the parameters needed to initialize Wiener deconvolution ,
[0034] in, is the homomorphic encryption model gradient obtained after server-side aggregation, For system response, and The parameters obtained after Fourier transform;
[0035] 6.2 Calculate the power spectral density of the signal based on the parameters after Fourier transformation: ,
[0036] 6.3 Calculating the initial Wiener deconvolution based on power spectral density and other parameters ,
[0037] in, for The complex conjugate of is the power spectral density of the noise,
[0038] 6.4 Initial filtering of model parameters based on Wiener deconvolution ,
[0039] 6.5 The modified power spectrum is calculated based on the power spectrum in the previous iteration.
[0040] ,
[0041] 6.6 Update the power spectrum based on the corrected power spectrum and the previous round of filtering results.
[0042] ,
[0043] 6.7 Update Wiener Deconvolution Based on the New Power Spectrum ,
[0044] 6.8 Apply the new Wiener deconvolution to filter until the result converges. The calculation formula is: ,
[0045] 6.9 Perform inverse Fourier transform on the filtered result ,
[0046] Among them, firstly, the modified power spectrum is calculated according to the power spectrum in the previous iteration , and then update the power spectrum based on the corrected power spectrum and the previous round of filtering results , and then get the Wiener deconvolution , and finally get the new filtering result , The server obtains the homomorphically encrypted model parameters obtained after the inverse Fourier transform and sends them to each client. Each client uses its own homomorphic encryption private key to decrypt them and then uses them for the next round of local model training.
[0047] An electronic device comprises a memory, a processor and a computer program stored in the memory and executable on the processor, wherein when the processor executes the program, a federated learning method based on Wiener deconvolution and diagonal Fisher pruning is implemented.
[0048] A computer-readable storage medium stores computer instructions, which, when executed by a processor, implement a federated learning method based on Wiener deconvolution and diagonal Fisher pruning.
[0049] Compared with the prior art, the present invention has the following advantages:
[0050] After completing the model training, the client of the present invention uses the Diag-Fisher pruning mechanism to prune the gradient through the Fisher information matrix, retains key information through the Fisher information matrix, optimizes the importance ranking of model parameter transmission from a statistical perspective, effectively reduces communication overhead, and avoids significant impact on model performance. Then, combined with the natural gradient optimization method, this method is based on the Fisher information matrix and adjusts the optimization direction to adapt to the geometric structure of the parameter space, which not only speeds up the convergence speed of the model, but also improves the final training accuracy. After the server receives the encrypted pruned gradient, a denoising strategy based on Wiener deconvolution is adopted. The deconvolution filter is designed by estimating the noise power spectral density, which minimizes the interference of noise on the effective gradient information from the signal processing perspective, and further enhances the accuracy of the model. BRIEF DESCRIPTION OF THE DRAWINGS
[0051] Figure 1 It is a framework diagram of a federated learning method based on Wiener deconvolution and diagonal Fisher pruning of the present invention;
[0052] Figure 2 It is a flow chart of pruning based on the Diag-Fisher pruning mechanism and using natural gradient for optimization in the present invention;
[0053] Figure 3 It is a flow chart of filtering based on Wiener deconvolution in the present invention. DETAILED DESCRIPTION
[0054] In order to deepen the understanding of the present invention, the present embodiment is described in detail below with reference to the accompanying drawings.
[0055] Example: See Figure 1-Figure 3 , a federated learning method based on Wiener deconvolution and diagonal Fisher pruning, comprising the following steps:
[0056] Assume that there is a server and n clients in a video recommendation system. Client 1, client 2 and client 3 are any three clients connected to the server and capable of communicating. The clients and the server jointly train the model. This embodiment takes client 1, client 2, client 3 and the server as examples to give a detailed description of the training process. The entire federated learning architecture is as follows: Figure 1 As shown in the figure, it mainly includes two processes: Diag-Fisher pruning and Wiener deconvolution denoising.
[0057] Step 1: Each user end of the video recommendation system receives the initial global model from the server end and trains it based on local video browsing data.
[0058] Step 2: Each user terminal performs local training on the received global model based on the local data set to obtain a local model.
[0059] Step 3: Each user end of the video recommendation system uses the Diag-Fisher pruning mechanism to prune the gradient of the local model, such as Figure 2 As shown, the process is as follows:
[0060] 3.1 Initialize the diagonal Fisher information matrix corresponding to the local model, and its calculation expression is: ,
[0061] 3.2 For each batch of data , forward propagation calculates the loss ,
[0062] in, For Clients The local model gradient of A batch of data,
[0063] 3.3 For each batch of data , calculate the gradient based on the loss back propagation ,
[0064] in, For Clients The local model gradient of For its data set A batch of data, represents the loss calculated by forward propagation,
[0065] 3.4 For each batch of data , the diagonal Fisher information matrix of the local model is accumulated ,in, The gradient calculated by back propagation,
[0066] 3.5 According to the batch dataset that represents the current client for updating the model The size of the normalized diagonal Fisher information matrix ,in, Represents the batch dataset used by the current client to update the model The size of
[0067] 3.6 Calculate the pruning threshold based on the set pruning rate and the calculated diagonal Fisher information matrix ,
[0068] in, is the calculated diagonal Fisher information matrix, is the pre-set pruning rate, The pruning threshold is obtained by calculation.
[0069] 3.7 Compare the Fisher information value in the diagonal Fisher information matrix with the pruning threshold to generate a mask
[0070] ,
[0071] 3.8 Apply mask to prune local model parameters: ,in, The pruning threshold is obtained by calculation. The mask is calculated, and finally the mask is applied to get the new gradient after pruning ,
[0072] Step 4: Use natural gradient optimization on the pruned local model. The process is as follows: Figure 2 ,
[0073] 4.1 For each batch of data , forward propagation calculates the loss ,
[0074] 4.2 For each batch of data , calculate the gradient based on the loss back propagation
[0075] ,
[0076] 4.3 Obtaining the natural gradient based on the diagonal Fisher information matrix and gradient calculation ,
[0077] in, For Clients The pruned local model gradients, is a small positive value, which is a smoothing term used to prevent the denominator from being zero and ensure numerical stability. is the calculated natural gradient,
[0078] 4.4 Update local model parameters according to the calculated natural gradient: ,
[0079] in, is the learning rate, is the local model gradient after natural gradient update,
[0080] The client uses natural gradient optimization on the pruned local model to optimize and correct the pruned gradient to ensure that the gradient update direction conforms to the geometric structure of the parameter space, thereby improving the convergence speed and accuracy of the model.
[0081] Step 5: Each client uploads the pruned model parameters and optimized gradients to the server through homomorphic encryption.
[0082] Step 6: Use the federated mechanism based on Wiener deconvolution to denoise the encrypted gradient. The process is as follows: Figure 3 ,
[0083] 6.1 Fourier transform of the parameters needed to initialize Wiener deconvolution ,
[0084] in, is the homomorphic encryption model gradient obtained after server-side aggregation, For system response, and The parameters obtained after Fourier transform;
[0085] 6.2 Calculate the power spectral density of the signal based on the parameters after Fourier transformation: ,
[0086] 6.3 Calculating the initial Wiener deconvolution based on power spectral density and other parameters ,
[0087] in, for The complex conjugate of is the power spectral density of the noise,
[0088] 6.4 Initial filtering of model parameters based on Wiener deconvolution ,
[0089] 6.5 Calculate the corrected power spectrum based on the power spectrum in the previous iteration
[0090] ,
[0091] 6.6 Update the power spectrum based on the corrected power spectrum and the previous round of filtering results
[0092] ,
[0093] 6.7 Update Wiener Deconvolution Based on the New Power Spectrum ,
[0094] 6.8 Apply the new Wiener deconvolution to filter until the result converges. The calculation formula is: ,
[0095] 6.9 Perform inverse Fourier transform on the filtered result .
[0096] Among them, the model parameters are homomorphically encrypted after inverse Fourier transform. After the server obtains them, it sends them to each user end of the video recommendation system. Each user end uses its own homomorphic encryption private key to decrypt them, and then uses them for the next round of local model training.
[0097] It should be noted that the above embodiments are not intended to limit the protection scope of the present invention, and equivalent changes or substitutions made on the basis of the above technical solutions all fall within the protection scope of the claims of the present invention.
Claims
1. A federated learning method based on Wiener deconvolution and diagonal Fisher pruning, characterized in that: The method comprises the following steps: Step 1: The client receives the global model from the server and trains it based on the local dataset. Step 2: The client performs local training on the received global model based on the local data set to obtain a local model. Step 3: Use the Diag-Fisher pruning mechanism to prune the gradient of the local model. Step 4: Use natural gradient optimization on the pruned local model. Step 5: The client uploads the pruned model parameters and optimized gradients to the server through homomorphic encryption. Step 6: Use the federation mechanism based on Wiener deconvolution to denoise the encrypted gradients. The server aggregates the denoised gradients from multiple clients to obtain the aggregated global model. The server sends the updated global model to the client to complete the iteration of the current training round. When the preset iteration round is reached, the server finally generates a global model trained by federated learning. Step 3: Use the Diag-Fisher pruning mechanism to prune the gradient of the local model. The process is as follows: 3.1 Initialize the diagonal Fisher information matrix corresponding to the local model, which is calculated as: F diag ←0, 3.2 For each batch of data (x, y) ∈ D i , forward propagation calculates the loss L(θ i , x, y), Among them, θ i is the local model gradient of client i, and its dataset D i A batch of data, 3.3 For each batch of data (x, y) ∈ D i , calculate the gradient based on the loss back propagation Among them, θ i is the local model gradient of client i, (x, y) is its dataset D i A batch of data, L(θ i , x, y) represents the loss calculated by forward propagation, 3.4 For each batch of data (x, y) ∈ D i , the diagonal Fisher information matrix F of the local model is accumulated diag ←F diag +g 2 , where g is the gradient calculated by back propagation. 3.5 According to the batch dataset D that the current client uses to update the model i The size of the normalized diagonal Fisher information matrix F diag ←F diag / len(D i ),in, len(D i ) represents the batch dataset D used by the current client to update the model i The size of 3.6 According to the set pruning rate and the calculated diagonal Fisher information matrix, calculate the pruning threshold τ←quantile(F diag ,p), in, F diag is the calculated diagonal Fisher information matrix, p is the pre-set pruning rate, τ is the calculated pruning threshold, 3.7 Compare the Fisher information value in the diagonal Fisher information matrix with the pruning threshold to generate the mask M←F diag >τ, 3.8 Applying masks to prune local model parameters: θ′ i ←θ i ⊙M, where τ is the calculated pruning threshold, M is the calculated mask, and finally the mask is applied to obtain the new gradient θ′ after pruning i .
2. According to the federated learning method based on Wiener deconvolution and diagonal Fisher pruning in claim 1, it is characterized in that, in step 4, natural gradient optimization is used on the pruned local model, and the process is as follows: 4.1 For each batch of data (x, y)∈D i , forward propagation calculates the loss L(θ′ i ,x,y), 4.2 For each batch of data (x,y)∈D i , calculate the gradient based on the loss back propagation 4.3 Obtaining the natural gradient based on the diagonal Fisher information matrix and gradient calculation in, θ′ i is the local model gradient of client i after pruning, ∈ is a small positive value, which is a smoothing term used to prevent the denominator from being zero and ensure numerical stability. is the calculated natural gradient, 4.4 Update local model parameters according to the calculated natural gradient: ,in, η is the learning rate, θ″ i is the local model gradient after natural gradient update.
3. The federated learning method based on Wiener deconvolution and diagonal Fisher pruning according to claim 1, characterized in that: Step 6: Use the federated mechanism based on Wiener deconvolution to denoise the encrypted gradient. The process is as follows: 6.1 Fourier transform of the parameters needed to initialize Wiener deconvolution Among them, E(θ″ i ) is the homomorphic encryption model gradient obtained after server-side aggregation, h is the system response, G and H are the parameters obtained after Fourier transform, 6.2 Calculate the power spectral density of the signal based on the parameters after Fourier transformation: P f (0)←|G| 2 , 6.3 Calculating the initial Wiener deconvolution based on power spectral density and other parameters in, H * is the conjugate complex number of H, P n is the power spectral density of the noise, 6.4 Perform initial filtering of model parameters according to Wiener deconvolution F(0) = W(0)G, 6.5 Calculate the corrected power spectrum based on the power spectrum in the previous iteration 6.6 Update the power spectrum P based on the corrected power spectrum and the previous round of filtering results f (i)=|F(i-1)| 2 +P f (i)correction 6.7 Update Wiener Deconvolution Based on the New Power Spectrum 6.8 Apply the new Wiener deconvolution to filter until the result converges. The calculation formula is: F(i)=W(i)G, 6.9 Perform inverse Fourier transform on the filtered result f = IFFT(F(i)), in, First, the modified power spectrum P is calculated based on the power spectrum in the previous iteration. f (i) Correction, then update the power spectrum P according to the corrected power spectrum and the filtering result of the previous round f (i), and then obtain the Wiener deconvolution W(i), and finally obtain the new filtering result F(i), where f is the homomorphically encrypted model parameter obtained after the inverse Fourier transform. After the server obtains it, it sends it to each client. Each client uses its own homomorphic encryption private key to decrypt it, and then uses it for the next round of local model training.
4. An electronic device comprising a memory, a processor, and a computer program stored in the memory and executable on the processor, wherein: When the processor executes the program, a federated learning method based on Wiener deconvolution and diagonal Fisher pruning as described in any one of claims 1 to 3 is implemented.
5. A computer-readable storage medium having computer instructions stored thereon, characterized in that: When the computer instruction is executed by the processor, a federated learning method based on Wiener deconvolution and diagonal Fisher pruning as described in any one of claims 1 to 3 is implemented.
Citation Information
Patent Citations
Wireless channel scene classification method
CN106548136A
Efficient federated learning system and method based on channel pruning
CN117592538A
Model training method based on lossless federated learning and related equipment
CN118734940A