Model parameter update method, apparatus and non-volatile storage medium
Patent Information
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2026-04-27
- Publication Date
- 2026-08-11
AI Technical Summary
[0004]本发明实施例提供了一种模型参数更新方法、装置和非易失性存储介质,以至少解决由于目前采用医疗数据进行模型训练时需获取每个客户端的医疗数据造成的数据泄漏的技术问题
[0015] In this embodiment of the invention, a model parameter update method is adopted. This involves obtaining the objective function parameters corresponding to each of multiple clients; receiving the representation matrices corresponding to each of the multiple clients; calculating the similarity value between any two clients based on the representation matrices; determining the similarity value between the multiple clients; determining the target global model parameters based on the similarity value between the multiple clients and the objective function parameters corresponding to each client; and distributing the target global model parameters to the multiple clients. This achieves the goal of protecting the privacy and security of local data, thereby improving the interpretability, convergence speed, and accuracy of federated learning, and enhancing its stability. Furthermore, it solves the technical problem of data leakage caused by the need to obtain medical data from each client when using medical data for model training.
Smart Images

Figure CN122549537A_ABST
Abstract
Description
Technical Field
[0001] This invention relates to the fields of federated learning and information security technology, and more specifically, to a model parameter update method, apparatus, and non-volatile storage medium. Background Technology
[0002] Collecting and managing large-scale medical datasets from multiple institutions is crucial for training accurate deep learning models, but privacy concerns often hinder data sharing. Currently, Artificial Intelligence (AI), as the most popular research and application area, has fully entered the big data era centered on deep learning. Machine learning based on big data has not only driven the rapid development of AI but has also raised a series of information security risks. These risks stem from the learning mechanism of deep learning itself; whether in the model training phase, or in the model inference and application phases, there is a risk of data leakage or malicious exploitation, which, if it occurs, will have serious consequences.
[0003] There is currently no effective solution to the above problems. Summary of the Invention
[0004] This invention provides a model parameter update method, apparatus, and non-volatile storage medium to at least solve the technical problem of data leakage caused by the need to obtain medical data from each client when using medical data for model training.
[0005] According to one aspect of the present invention, a model parameter update method is provided, comprising: obtaining objective function parameters corresponding to each of multiple clients; receiving representation matrices corresponding to each of the multiple clients; calculating a similarity value between any two clients among the multiple clients based on the representation matrices corresponding to each client, and determining a similarity value between the multiple clients; determining a target global model parameter based on the similarity value between the multiple clients and the objective function parameters corresponding to each client; and distributing the target global model parameter to the multiple clients.
[0006] Optionally, receiving the representation matrices corresponding to each of the multiple clients includes: sending sample medical datasets to multiple clients, wherein each client performs forward propagation on the sample medical datasets based on its own local neural network model to obtain the representation matrices corresponding to each client; and receiving the respective representation matrices fed back by the multiple clients.
[0007] Optionally, based on the representation matrices corresponding to each of the multiple clients, a similarity value between the first client and the second client among the multiple clients is calculated, wherein the first client and the second client are any two clients among the multiple clients, including: determining the sample point sets corresponding to the first client and the second client respectively according to the representation matrices corresponding to the first client and the second client respectively, wherein a row vector in the representation matrix represents a sample point; calculating the Jensen-Shannon divergence between the first client and the second client based on the sample point sets corresponding to the first client and the second client respectively; and determining the similarity value between the first client and the second client based on the Jensen-Shannon divergence between the first client and the second client.
[0008] Optionally, the target global model parameters are determined based on the similarity values between multiple clients and the objective function parameters corresponding to each client, including: determining the aggregate weights corresponding to each client based on the similarity values between multiple clients; determining the initial global model parameters based on the aggregate weights corresponding to each client and the objective function parameters corresponding to each client; and updating the initial global model parameters using the gradient descent algorithm to obtain the target global model parameters.
[0009] Optionally, the initial global model parameters are determined based on the aggregate weights and objective function parameters corresponding to each of the multiple clients, including: identifying clients whose aggregate weights exceed a preset weight threshold as relevant clients; and determining the initial global model parameters based on the aggregate weights and objective function parameters corresponding to the relevant clients.
[0010] Optionally, it also includes: repeatedly determining the target global model parameters and sending the target global model parameters to multiple clients until the accuracy of the local neural network model corresponding to the target client in the multiple clients reaches a preset accuracy threshold, then stopping the repetition and determining the current target global model parameters as the target model parameters corresponding to the target client.
[0011] According to another aspect of the present invention, a model parameter update apparatus is also provided, comprising: an acquisition module for acquiring objective function parameters corresponding to each of multiple clients; a receiving module for receiving representation matrices corresponding to each of the multiple clients; a calculation module for calculating a similarity value between any two clients among the multiple clients based on the representation matrices corresponding to each client, and determining a similarity value between the multiple clients; a determination module for determining target global model parameters based on the similarity value between the multiple clients and the objective function parameters corresponding to each client; and a sending module for sending the target global model parameters to the multiple clients.
[0012] According to another aspect of the present invention, a non-volatile storage medium is also provided, the non-volatile storage medium including a stored program, wherein, when the program is running, the device where the non-volatile storage medium is located is controlled to execute any of the above-described model parameter update methods.
[0013] According to another aspect of the present invention, a computer device is also provided, the computer device including a processor, the processor being configured to run a program, wherein the program executes any of the model parameter update methods described above during runtime.
[0014] According to another aspect of the present invention, a computer program product is also provided, including a computer program that, when executed by a processor, implements any of the above-described model parameter update methods.
[0015] In this embodiment of the invention, a model parameter update method is adopted. This involves obtaining the objective function parameters corresponding to each of multiple clients; receiving the representation matrices corresponding to each of the multiple clients; calculating the similarity value between any two clients based on the representation matrices; determining the similarity value between the multiple clients; determining the target global model parameters based on the similarity value between the multiple clients and the objective function parameters corresponding to each client; and distributing the target global model parameters to the multiple clients. This achieves the goal of protecting the privacy and security of local data, thereby improving the interpretability, convergence speed, and accuracy of federated learning, and enhancing its stability. Furthermore, it solves the technical problem of data leakage caused by the need to obtain medical data from each client when using medical data for model training. Attached Figure Description
[0016] The accompanying drawings, which are included to provide a further understanding of the invention and form part of this application, illustrate exemplary embodiments of the invention and, together with their description, serve to explain the invention and do not constitute an undue limitation thereof. In the drawings:
[0017] Figure 1 A hardware structure block diagram of a computer terminal for implementing a model parameter update method is shown.
[0018] Figure 2 This is a flowchart illustrating the model parameter update method provided in an embodiment of the present invention;
[0019] Figure 3 This is a structural block diagram of a model parameter updating device provided according to an embodiment of the present invention. Detailed Implementation
[0020] To enable those skilled in the art to better understand the present invention, the technical solutions of the present invention will be clearly and completely described below with reference to the accompanying drawings of the embodiments of the present invention. Obviously, the described embodiments are only some embodiments of the present invention, and not all embodiments. Based on the embodiments of the present invention, all other embodiments obtained by those skilled in the art without creative effort should fall within the scope of protection of the present invention.
[0021] It should be noted that the terms "first," "second," etc., in the specification, claims, and accompanying drawings of this invention are used to distinguish similar objects and are not necessarily used to describe a specific order or sequence. It should be understood that such data can be interchanged where appropriate so that the embodiments of the invention described herein can be implemented in orders other than those illustrated or described herein. Furthermore, the terms "comprising" and "having," and any variations thereof, are intended to cover a non-exclusive inclusion; for example, a process, method, system, product, or apparatus that comprises a series of steps or units is not necessarily limited to those steps or units explicitly listed, but may include other steps or units not explicitly listed or inherent to such processes, methods, products, or apparatus.
[0022] According to an embodiment of the present invention, a method embodiment for updating model parameters is provided. It should be noted that the steps shown in the flowchart in the accompanying drawings can be executed in a computer system such as a set of computer-executable instructions. Furthermore, although a logical order is shown in the flowchart, in some cases, the steps shown or described may be executed in a different order than that shown here.
[0023] The method embodiment provided in Embodiment 1 of this application can be executed on a mobile terminal, computer terminal, or similar computing device. Figure 1 A hardware block diagram of a computer terminal for implementing a model parameter update method is shown. Figure 1 As shown, the computer terminal 10 may include one or more processors (shown as 102a, 102b, ..., 102n in the figure) (the processor may include, but is not limited to, a microprocessor MCU or a programmable logic device FPGA, etc.) and a memory 104 for storing data. In addition, it may also include: a display, an input / output interface (I / O interface), a universal serial bus (USB) port (which may be included as one of the ports of a BUS bus), a network interface, a power supply, and / or a camera. Those skilled in the art will understand that... Figure 1 The structure shown is for illustrative purposes only and does not limit the structure of the aforementioned electronic device. For example, computer terminal 10 may also include... Figure 1 The more or fewer components shown, or having the same Figure 1The different configurations shown.
[0024] It should be noted that the aforementioned one or more processors and / or other data processing circuits are generally referred to herein as "data processing circuits". These data processing circuits may be embodied, in whole or in part, in software, hardware, firmware, or any other combination thereof. Furthermore, the data processing circuits may be a single, independent processing module, or may be integrated, in whole or in part, into any other element within the computer terminal 10. As involved in the embodiments of this application, the data processing circuits serve as a processor control mechanism (e.g., selection of a variable resistor termination path connected to an interface).
[0025] The memory 104 can be used to store software programs and modules of application software, such as the program instructions / data storage device corresponding to the model parameter update method in this embodiment of the invention. The processor executes various functional applications and data processing by running the software programs and modules stored in the memory 104, thereby implementing the above-mentioned model parameter update method of the application. The memory 104 may include high-speed random access memory, and may also include non-volatile memory, such as one or more magnetic storage devices, flash memory, or other non-volatile solid-state memory. In some instances, the memory 104 may further include memory remotely located relative to the processor, and these remote memories can be connected to the computer terminal 10 via a network. Examples of such networks include, but are not limited to, the Internet, corporate intranets, local area networks, mobile communication networks, and combinations thereof.
[0026] The display can be, for example, a touchscreen liquid crystal display (LCD) that allows the user to interact with the user interface of the computer terminal 10.
[0027] Figure 2 This is a flowchart illustrating the model parameter update method provided in an embodiment of the present invention, as shown below. Figure 2 As shown, the method includes the following steps:
[0028] Step S202: Obtain the target function parameters corresponding to each of the multiple clients.
[0029] Before federated learning can begin, the server needs to create a federated learning environment. This includes defining the model architecture for training, initializing the model parameters, and setting the hyperparameters of federated learning, such as the learning rate and the number of iterations.
[0030] Furthermore, the server distributes the initialized model parameters to all clients participating in federated learning. These clients may be data analysis nodes from different hospitals, research institutions, or companies. They each possess their own private datasets but do not share this data. After receiving the initial model parameters, each client trains its model on its local dataset. This training process is based on the client's own data, thus generating model parameter updates adapted to the characteristics of the local data. The objective function parameters include minimizing the model's loss function on the local data, which is the primary objective for the client to optimize its model parameters. In this embodiment, this loss function is related to the accuracy of disease prediction and the precision of image recognition. After local training is complete, the client calculates its model's objective function parameters on its local dataset. This includes not only the model parameters themselves but also performance metrics on the local data, such as loss value and accuracy, for subsequent model evaluation and aggregation.
[0031] Preferably, the objective function of this step as follows:
[0032] ;
[0033] ;
[0034] in, This is a very long vector formed by concatenating the model parameters from the global parameter vector for all clients. This is the transpose of the global parameter vector. For the overall loss, For Laplace regularization, for The weight of the Laplace regularization term, For Laplace matrix, Indicates the current client. Indicates the total number of clients. Indicates other clients, Represents a collection of other clients. and This represents two neural network parameters. The aggregation strength of these two neural network parameters, This represents the square of the Euclidean norm.
[0035] Finally, the client reports the calculated objective function parameters to the server, without reporting the specific local data, thus ensuring that medical data is not leaked. The server, acting as the coordinator of federated learning, is responsible for receiving the objective function parameters from each client. These parameters will be used in the next step of model aggregation and optimization to form a global model.
[0036] Step S204: Receive the representation matrices corresponding to each of the multiple clients.
[0037] In this step, the server first prepares a public sample medical dataset. This dataset is used to evaluate and compare model performance on different clients. The sample medical dataset is sufficiently diverse, but its size is also carefully considered to reduce communication overhead. The server then distributes this sample medical dataset to each client participating in federated learning, ensuring that each client receives the same dataset.
[0038] Upon receiving the sample medical dataset, each client uses its local neural network model to perform forward propagation on the dataset, generating a representation matrix. This representation matrix is a set of activation values from a specific layer (usually a hidden layer), where each row vector represents the representation of a sample point in the neural network. Assume that client i receives the sample medical dataset as follows: , where n is the number of samples. The neural network model parameters for client i are: Then the generated representation matrix Can be regarded as A feature matrix, where d represents the feature dimension. The specific calculation is as follows:
[0039] ,
[0040] Where f is the forward propagation function of the neural network, Each column vector in This represents the representation of the k-th sample point under the client i model.
[0041] Specifically, each client will send their representation matrix Feedback is sent back to the server. To protect data privacy, the server does not request the client to provide the original medical data. Instead, it requires the client to perform forward propagation locally, transmitting only the representation matrix.
[0042] Step S206: Based on the representation matrices corresponding to each of the multiple clients, calculate the similarity value between any two clients among the multiple clients, and determine the similarity value between the multiple clients.
[0043] In this step, for each client, each row of the received representation matrix is treated as a sample point. For example, for client i, its representation matrix is... ,in The number of rows represents the number of samples, and the number of columns represents the feature dimension. In this way, we can... Convert to a by A set of sample points ,in yes The number of rows. The constructed set of sample points. This is considered as the empirical distribution of client i on a public dataset. The purpose of this step is to transform the high-dimensional feature representation into a comparable distribution form, thereby enabling the application of statistical measures for comparison.
[0044] Furthermore, for any two clients i and j, calculate their empirical distributions. and The Jensen-Shannon (JS) divergence between two probability distributions. The JS divergence is a statistic that measures the difference between two probability distributions and is commonly used to compare the similarity of two sets. Its calculation formula is as follows:
[0045] ,
[0046] in, KL Indicates the Kullback-Leibler divergence. M yes and The average distribution of . By calculating the JS divergence, a non-negative value can be obtained. The smaller the value, the closer the empirical distributions of the two clients are. To transform the JS divergence into a more intuitive similarity value, and restrict it to the interval ([0,1]), the following formula can also be used for calculation:
[0047] ,
[0048] in, That is, the client i and j The similarity value between clients i and j. The closer the value is to 1, the closer the representation distributions of clients i and j are, meaning the more similar the feature representations of the two clients are; conversely, the closer the value is to 0, the greater the difference is.
[0049] Finally, repeat the above process to calculate the pairwise similarity values between all clients, forming a similarity matrix. This matrix will guide the subsequent parameter aggregation; that is, the aggregation strength between clients will be determined based on their similarity. Clients with higher similarity will have greater weight in parameter aggregation. After determining the similarity matrix, the similarity values can be further mapped to aggregation weights, i.e., as edge weights in the Laplace regularization term. This mapping process can be achieved by using a regularization strength coefficient and a nonlinear amplification factor.
[0050] Step S208: Determine the target global model parameters based on the similarity values between multiple clients and the objective function parameters corresponding to each client.
[0051] In this step, the aggregation weight of each client when aggregating with other clients is first determined based on the calculated similarity values between clients. The aggregation weight reflects the importance of a client in the aggregation process and is usually proportional to the similarity value; that is, a client with a higher similarity value will have a greater weight during aggregation. The aggregation weight can be determined using the following formula:
[0052] ,
[0053] in, This represents the aggregate weight between client i and client j. It is the similarity value between client i and client j. It is a non-linear amplification factor, and N is the total number of clients.
[0054] After determining the aggregation weights among the various clients, the initial global model parameters are calculated based on these weights and the objective function parameters of each client. This can be achieved by weighted averaging of the model parameters for each client:
[0055] ,
[0056] in, This represents the initial global model parameters in round t+1. Let the parameters of the objective function for client i in round t be represented. It is the aggregate weight of client i. This weight can be the weighted average of the aggregate weights of client i and all other clients, or it can be the weighted average of the aggregate weights of client i and a specific set of related clients.
[0057] Finally, gradient descent is used to optimize the initial global model parameters to further improve model performance. This optimization process may involve multiple iterations until certain termination conditions are met, such as the global loss function reaching a preset threshold or the model's accuracy no longer significantly improving. Preferably, the optimization formula is as follows:
[0058] ,
[0059] in, This represents the target global model parameters after being updated using the gradient descent algorithm. It's the learning rate. Indicates the global loss function in The gradient at that point.
[0060] Step S210: Distribute the target global model parameters to multiple clients.
[0061] In this step, the determined target global model parameters are first packaged, ensuring that all necessary parameters are correctly encapsulated. To facilitate transmission over the network, they also need to be encoded, such as using JSON or protobuf formats, to reduce transmission size and improve efficiency. Considering data security and privacy protection in a federated learning environment, encrypting the packaged global model parameters is necessary. Public-key encryption techniques, such as RSA or elliptic curve cryptography, can be used to ensure that only the corresponding client can decrypt and use these parameters.
[0062] Furthermore, the encrypted global model parameters are distributed to all participating clients through established communication channels. This relies on network infrastructure, such as the internet or a dedicated network, to ensure the parameters arrive at their destination securely and promptly. In practice, multicast or broadcast techniques can be used to send the parameters to all clients at once, thus saving network resources. Each client, upon receiving the encrypted global model parameters, decrypts them using its own private key. Upon successful decryption, the client updates its local neural network model using the new global model parameters. This process may involve overwriting old model parameters or performing parameter fusion according to specific update rules.
[0063] To ensure correct transmission and reception of parameters, the client sends an acknowledgment message to the server after updating its local model, indicating that the parameter update is complete. The server can then use this feedback to assess the success rate of parameter distribution and take appropriate remedial measures, such as resending parameters to clients that have not yet acknowledged receipt.
[0064] The above process is repeated continuously as the federated learning training process iterates. The server continuously receives new parameters and representation matrices from the client, recalculates the global model parameters after each iteration, and sends them back to the client until a pre-set stopping condition is met, such as reaching a certain accuracy threshold or the number of iterations.
[0065] Through the above steps, the interpretability, convergence speed, and accuracy of federated learning can be improved, and its stability can be enhanced. This solves the technical problem of data leakage caused by the need to obtain medical data from each client when using medical data for model training.
[0066] As an optional embodiment, this can be achieved through the following steps: receiving the representation matrices corresponding to multiple clients, including: sending sample medical datasets to multiple clients, wherein each client performs forward propagation on the sample medical datasets based on its corresponding local neural network model to obtain the representation matrices corresponding to each client; and receiving the corresponding representation matrices fed back by multiple clients.
[0067] In this step, the central server distributes a public sample medical dataset to all clients participating in federated learning. This dataset typically contains insensitive samples used to measure model performance, such as a standardized set of medical images or health record summaries. Each client, upon receiving the public sample medical dataset, feeds it into its local neural network model for forward propagation to generate a representation matrix. Assume that client i receives the sample medical dataset as follows: , where n is the number of samples. The neural network model parameters for client i are: Then the generated representation matrix Can be regarded as A feature matrix, where d represents the feature dimension. The specific calculation is as follows:
[0068] ,
[0069] Where f is the forward propagation function of the neural network, Each column vector in This represents the representation of the k-th sample point under the client i model.
[0070] Specifically, each client will send their representation matrix Feedback is sent back to the server. To protect data privacy, the server does not request the client to provide the original medical data. Instead, it requires the client to perform forward propagation locally, transmitting only the representation matrix.
[0071] As an optional embodiment, this can be achieved through the following steps: Calculating the similarity value between a first client and a second client among the multiple clients, based on the representation matrices corresponding to each of the multiple clients, wherein the first client and the second client are any two clients among the multiple clients, including: determining the sample point sets corresponding to the first client and the second client respectively according to their respective representation matrices, wherein a row vector in the representation matrix represents a sample point; calculating the Jensen-Shannon divergence between the first client and the second client based on their respective sample point sets; and determining the similarity value between the first client and the second client based on the Jensen-Shannon divergence.
[0072] In this step, the representation matrices of the first and second clients are first obtained. These representation matrices are obtained by forward propagating a common sample medical dataset onto the local neural network model of each client, where each row vector represents the hidden layer feature representation of a sample point in the model. For the first client, assume the representation matrix is... It is made of It consists of ___ sample points, each with D-dimensional features. Similarly, at the second client, the representation matrix is ___. It is made of It consists of 100 sample points, and each sample point also has D-dimensional features. These matrices are then transformed into a set of sample points, i.e. and This is so that the distribution similarity can be calculated later.
[0073] Furthermore, Jensen-Shannon divergence (JS divergence) is used to measure the distributional difference in feature representations between the first and second clients. Considering... and Having approximated the empirical distributions of the first and second clients on the public dataset, we first need to estimate the probability distributions of these sets. This is typically done through kernel density estimation, treating each sample point as a kernel function (such as a Gaussian kernel) placed at its corresponding location. Given the kernel bandwidth h, the estimated distributions of the first and second clients are:
[0074] ;
[0075] ;
[0076] Where K is the Gaussian kernel function, Representation matrix The i-th row vector, Representation matrix The j-th row vector. Further, define the mixture distribution. as follows:
[0077] ;
[0078] Next, calculate the JS divergence between the first client and the second client:
[0079] ;
[0080] Finally, the JS divergence needs to be converted into similarity values, ensuring the values fall within the range [0,1] for easy use in subsequent model aggregation. The following conversion formula is used:
[0081] ;
[0082] From this, we obtain The closer the value is to 1, the closer the client representation distributions are, and the more similar the two are; the closer the value is to 0, the greater the difference is.
[0083] As an optional embodiment, this can be achieved through the following steps: determining the target global model parameters based on the similarity values between multiple clients and the objective function parameters corresponding to each client, including: determining the aggregate weights corresponding to each client based on the similarity values between multiple clients; determining the initial global model parameters based on the aggregate weights corresponding to each client and the objective function parameters corresponding to each client; and updating the initial global model parameters using the gradient descent algorithm to obtain the target global model parameters.
[0084] In this step, the first step is to determine the aggregate weight for each client based on the previously calculated similarity values between clients. An aggregate weight function needs to be designed to convert the similarity values into aggregate weights. For example, a simple linear or non-linear mapping, such as an exponential function, a sigmoid function, or other functions, can be used to amplify the similarity values, thereby obtaining the aggregate weight for each client. In practical applications, the choice of aggregate weight function should consider the specific scenario and requirements of federated learning to ensure the rationality and fairness of weight allocation.
[0085] Based on the aggregate weights of each client and their respective objective function parameters, the initial global model parameters are calculated. This can be achieved through a weighted average, as shown in the following formula:
[0086]
[0087] Preferably, the initial global model parameters can be calculated. As edge weights in the Graph Laplace regularization term:
[0088]
[0089] in The normal intensity coefficient, As a nonlinear amplification factor, the gradient descent algorithm is then used to update and obtain the global parameters on the server side, and these parameters are used as the initial values for the neural network coefficients of each client in the next round.
[0090] Furthermore, after obtaining the initial global model parameters, the server uses the gradient descent algorithm to optimize and update these parameters, thereby further improving the performance of the global model. Gradient descent is an iterative optimization method that calculates the gradient of the objective function (global loss function) with respect to the model parameters and updates the parameters in the opposite direction of the gradient to find the optimal or near-optimal solution. Based on the above-mentioned optimized formula for calculating the initial global model parameters, the update can be performed according to the following formula:
[0091] ;
[0092] in Represents the client in round t+1. The updated model parameters; The assignment operator indicates that the new value calculated on the right is assigned to the variable on the left. Indicates the version of parameters used in the current round of regularization calculation; Represents the client in round t. The updated model parameters; Indicates the adjacent clients in round t. The updated model parameters; Indicates the weighting coefficient; Indicates the learning rate; for and The strength of aggregation.
[0093] As an optional embodiment, this can be achieved through the following steps: determining the initial global model parameters based on the aggregate weights corresponding to multiple clients and the objective function parameters corresponding to multiple clients, including: identifying clients whose aggregate weights exceed a preset weight threshold as relevant clients; and determining the initial global model parameters based on the aggregate weights corresponding to the relevant clients and the objective function parameters corresponding to the relevant clients.
[0094] In this step, the server or coordinator needs to define a preset weight threshold, which is used to filter out clients that have a high similarity to the global model. Specifically, the server checks the aggregate weights of all clients and marks those clients with weight values higher than the preset threshold as "relevant clients." The selection of the preset weight threshold needs to take into account the overall goals of federated learning, the heterogeneity of data distribution, and the need for privacy protection.
[0095] Once the relevant clients are identified, the server calculates a weighted average based on the aggregate weights of these clients. The aggregate weights reflect the similarity between the client's model and the global model; a larger weight indicates a greater contribution of that client's model parameters to the global model. Therefore, the server multiplies the objective function parameters (i.e., model coefficients) of each relevant client by its corresponding aggregate weight, then sums these products to obtain a total. This total is essentially a weighted merging of the model parameters from different clients, ensuring that the global model more accurately reflects the local model features that positively impact model performance. To ensure that the initial global model parameters maintain consistency with the previous model parameters, the server normalizes the obtained weighted average. Specifically, the server divides the obtained total by the sum of the aggregate weights of all relevant clients, thus obtaining a standardized initial global model parameter.
[0096] Finally, the server updates the global model using the initial global model parameters calculated above. This update can be viewed as an initialization of the global model parameters, setting them to the weighted average of the parameters of all relevant client models. In this way, the global model incorporates valuable information from multiple clients in the new round of training, helping to improve the model's accuracy and generalization ability.
[0097] As an optional embodiment, this can be achieved through the following steps: It further includes: repeatedly determining the target global model parameters and sending the target global model parameters to multiple clients until the accuracy of the local neural network model corresponding to the target client in the multiple clients reaches a preset accuracy threshold, then stopping the repetition and determining the current target global model parameters as the target model parameters corresponding to the target client.
[0098] After completing one round of training and model parameter updates, the server distributes the updated target global model parameters to all clients. Each client uses these global parameters to initialize its local neural network model and then continues local training, updating its own local model parameters. This process is called one iteration cycle of federated learning. Once all clients have completed local training and updated their local model parameters, they recalculate their representation matrices on the common sample set and feed them back to the server. Upon receiving the new representation matrices, the server recalculates the similarity values between clients and updates the client's aggregate weights accordingly. Based on the new weights and the clients' objective function parameters, the server uses gradient descent to calculate the new target global model parameters.
[0099] Furthermore, the server sends the newly obtained global model parameters back to the client, which uses this as the starting point for the next iteration, repeating the above process. This iterative training continues until a preset termination condition is reached. The preset termination condition is usually based on one of two situations: 1) The global loss function value is lower than a set threshold: In federated learning, the global loss function is used to measure the difference between the predictions of all clients and the true labels. When the global loss function value is low enough, meaning the model has learned enough information to accurately predict the results, the training process can be considered converged, and iteration can stop; 2) The accuracy of the local client model no longer improves: In some cases, even if the global loss function value is still higher than the threshold, the accuracy of the client model may no longer improve significantly. This usually means that the model has reached its performance bottleneck on the current dataset, and further training is unlikely to bring significant improvement. Therefore, when the accuracy of the local client model does not improve significantly in several consecutive iterations, it can also be regarded as a sign of training convergence.
[0100] It should be noted that, for the sake of simplicity, the foregoing method embodiments are all described as a series of actions. However, those skilled in the art should understand that the present invention is not limited to the described order of actions, because according to the present invention, some steps can be performed in other orders or simultaneously. Furthermore, those skilled in the art should also understand that the embodiments described in the specification are preferred embodiments, and the actions and modules involved are not necessarily essential to the present invention.
[0101] Through the above description of the embodiments, those skilled in the art can clearly understand that the model parameter update method according to the above embodiments can be implemented by means of software plus necessary general-purpose hardware platform. Of course, it can also be implemented by hardware, but in many cases the former is a better implementation method. Based on this understanding, the technical solution of the present invention, or the part that contributes to the prior art, can be embodied in the form of a software product. This computer software product is stored in a storage medium (such as ROM / RAM, magnetic disk, optical disk) and includes several instructions to cause a terminal device (which may be a mobile phone, computer, server, or network device, etc.) to execute the methods described in the various embodiments of the present invention.
[0102] According to an embodiment of the present invention, a model parameter updating apparatus for implementing the above-described model parameter updating method is also provided. Figure 3 This is a structural block diagram of a model parameter updating device provided according to an embodiment of the present invention, such as... Figure 3 As shown, the model parameter update device includes: an acquisition module 302, a receiving module 304, a calculation module 306, a determination module 308, and a sending module 310. The model parameter update device will be described below.
[0103] Module 302 is used to obtain the target function parameters corresponding to each of the multiple clients;
[0104] The receiving module 304, connected to the acquiring module 302, is used to receive the representation matrices corresponding to each of the multiple clients.
[0105] The calculation module 306, connected to the receiving module 304, is used to calculate the similarity value between any two clients among the multiple clients based on the representation matrix corresponding to each client, and to determine the similarity value between the multiple clients.
[0106] The determination module 308, connected to the calculation module 306, is used to determine the target global model parameters based on the similarity values between multiple clients and the objective function parameters corresponding to each client.
[0107] The sending module 310, connected to the determining module 308, is used to send the target global model parameters to multiple clients.
[0108] It should be noted that the acquisition module 302, receiving module 304, calculation module 306, determination module 308, and sending module 310 mentioned above correspond to steps S202 to S210 in the embodiments. Multiple modules and their corresponding steps implement the same instances and application scenarios, but are not limited to the content disclosed in the above embodiments. It should also be noted that the above modules, as part of the device, can run in the computer terminal 10 provided in the embodiments.
[0109] Embodiments of the present invention may provide a computer device. Optionally, in this embodiment, the computer device may be located in at least one of a plurality of network devices in a computer network. The computer device includes a memory and a processor.
[0110] The memory can be used to store software programs and modules, such as the program instructions / modules corresponding to the model parameter update method and apparatus in this embodiment of the invention. The processor executes various functional applications and data processing by running the software programs and modules stored in the memory, thereby realizing the aforementioned model parameter update method. The memory may include high-speed random access memory, and may also include non-volatile memory, such as one or more magnetic storage devices, flash memory, or other non-volatile solid-state memory. In some instances, the memory may further include memory remotely located relative to the processor, and these remote memories can be connected to a computer terminal via a network. Examples of such networks include, but are not limited to, the Internet, corporate intranets, local area networks, mobile communication networks, and combinations thereof.
[0111] The processor can access the information and application stored in the memory via the transmission device to perform the following steps: obtain the objective function parameters corresponding to each of the multiple clients; receive the representation matrix corresponding to each of the multiple clients; calculate the similarity value between any two clients based on the representation matrix corresponding to each client, and determine the similarity value between the multiple clients; determine the target global model parameters based on the similarity value between the multiple clients and the objective function parameters corresponding to each client; and send the target global model parameters to the multiple clients.
[0112] Optionally, the processor may also execute program code for the following steps: receiving representation matrices corresponding to multiple clients, including: sending sample medical datasets to multiple clients, wherein each client performs forward propagation on the sample medical dataset based on its corresponding local neural network model to obtain a representation matrix corresponding to each client; and receiving the corresponding representation matrices fed back by multiple clients.
[0113] Optionally, the processor may also execute program code for the following steps: calculating the similarity value between a first client and a second client among the multiple clients based on the representation matrices corresponding to each of the multiple clients, wherein the first client and the second client are any two clients among the multiple clients, including: determining the sample point set corresponding to each of the first client and the second client according to the representation matrices corresponding to each of the first client and the second client, wherein a row vector in the representation matrix represents a sample point; calculating the Jensen-Shannon divergence between the first client and the second client based on the sample point set corresponding to each of the first client and the second client; and determining the similarity value between the first client and the second client based on the Jensen-Shannon divergence between the first client and the second client.
[0114] Optionally, the processor may also execute program code that performs the following steps: determining the target global model parameters based on the similarity values between multiple clients and the objective function parameters corresponding to each client, including: determining the aggregate weights corresponding to each client based on the similarity values between multiple clients; determining the initial global model parameters based on the aggregate weights corresponding to each client and the objective function parameters corresponding to each client; and updating the initial global model parameters using a gradient descent algorithm to obtain the target global model parameters.
[0115] Optionally, the processor may also execute program code that performs the following steps: determining initial global model parameters based on the aggregate weights corresponding to multiple clients and the objective function parameters corresponding to multiple clients, including: determining clients whose aggregate weights exceed a preset weight threshold as relevant clients; and determining initial global model parameters based on the aggregate weights corresponding to relevant clients and the objective function parameters corresponding to relevant clients.
[0116] Optionally, the processor may also execute program code that includes the following steps: repeatedly determining the target global model parameters and sending the target global model parameters to multiple clients until the accuracy of the local neural network model corresponding to the target client in the multiple clients reaches a preset accuracy threshold, then stopping the repetition and determining the current target global model parameters as the target model parameters corresponding to the target client.
[0117] This invention provides a method for updating model parameters by acquiring the objective function parameters corresponding to multiple clients; receiving the representation matrices corresponding to multiple clients; calculating the similarity value between any two clients based on the representation matrices of each client, thus determining the similarity value between the multiple clients; determining the target global model parameters based on the similarity value between the multiple clients and the objective function parameters corresponding to each client; and distributing the target global model parameters to the multiple clients. This achieves the goal of protecting the privacy and security of local data, thereby improving the interpretability, convergence speed, and accuracy of federated learning, and enhancing its stability. Furthermore, it solves the technical problem of data leakage caused by the need to obtain medical data from each client when using medical data for model training.
[0118] Those skilled in the art will understand that all or part of the steps in the various methods of the above embodiments can be implemented by a program instructing the hardware related to the terminal device. The program can be stored in a non-volatile storage medium, which may include: flash drive, read-only memory (ROM), random access memory (RAM), disk or optical disk, etc.
[0119] Embodiments of the present invention also provide a non-volatile storage medium. Optionally, in this embodiment, the aforementioned non-volatile storage medium can be used to store the program code executed by the model parameter update method provided in the above embodiments.
[0120] Optionally, in this embodiment, the non-volatile storage medium may be located in any computer terminal in a group of computer terminals in a computer network, or in any mobile terminal in a group of mobile terminals.
[0121] Optionally, in this embodiment, the non-volatile storage medium is configured to store program code for performing the following steps: obtaining the objective function parameters corresponding to each of the multiple clients; receiving the representation matrix corresponding to each of the multiple clients; calculating the similarity value between any two clients among the multiple clients based on the representation matrix corresponding to each of the multiple clients, and determining the similarity value between the multiple clients; determining the target global model parameters based on the similarity value between the multiple clients and the objective function parameters corresponding to each of the multiple clients; and distributing the target global model parameters to the multiple clients.
[0122] Optionally, in this embodiment, the non-volatile storage medium is configured to store program code for performing the following steps: receiving representation matrices corresponding to multiple clients, including: sending sample medical datasets to multiple clients, wherein the multiple clients respectively perform forward propagation on the sample medical datasets based on their respective local neural network models to obtain representation matrices corresponding to each client; and receiving the respective representation matrices fed back by the multiple clients.
[0123] Optionally, in this embodiment, the non-volatile storage medium is configured to store program code for performing the following steps: calculating a similarity value between a first client and a second client among the multiple clients based on the representation matrices corresponding to each of the multiple clients, wherein the first client and the second client are any two clients among the multiple clients, including: determining a set of sample points corresponding to each of the first client and the second client according to the representation matrices corresponding to each of the first client and the second client, wherein a row vector in the representation matrix represents a sample point; calculating the Jensen-Shannon divergence between the first client and the second client based on the set of sample points corresponding to each of the first client and the second client; and determining a similarity value between the first client and the second client based on the Jensen-Shannon divergence between the first client and the second client.
[0124] Optionally, in this embodiment, the non-volatile storage medium is configured to store program code for performing the following steps: determining target global model parameters based on the similarity values between multiple clients and the objective function parameters corresponding to each client, including: determining the aggregate weights corresponding to each client based on the similarity values between multiple clients; determining the initial global model parameters based on the aggregate weights corresponding to each client and the objective function parameters corresponding to each client; and updating the initial global model parameters using a gradient descent algorithm to obtain the target global model parameters.
[0125] Optionally, in this embodiment, the non-volatile storage medium is configured to store program code for performing the following steps: determining initial global model parameters based on the aggregate weights corresponding to multiple clients and the objective function parameters corresponding to multiple clients, including: determining clients whose aggregate weights exceed a preset weight threshold as relevant clients; and determining initial global model parameters based on the aggregate weights corresponding to relevant clients and the objective function parameters corresponding to relevant clients.
[0126] Optionally, in this embodiment, the non-volatile storage medium is configured to store program code for performing the following steps: further including: repeatedly determining the target global model parameters and sending the target global model parameters to multiple clients until the accuracy of the local neural network model corresponding to the target client in the multiple clients reaches a preset accuracy threshold, then stopping the repetition, and determining the current target global model parameters as the target model parameters corresponding to the target client.
[0127] Embodiments of the present invention also provide a computer program product, including a computer program. Optionally, in this embodiment, when the computer program is executed by a processor, it can: obtain the objective function parameters corresponding to each of the multiple clients; receive the representation matrices corresponding to each of the multiple clients; calculate the similarity value between any two clients among the multiple clients based on the representation matrices corresponding to each client, and determine the similarity value between the multiple clients; determine the target global model parameters based on the similarity value between the multiple clients and the objective function parameters corresponding to each of the multiple clients; and distribute the target global model parameters to the multiple clients.
[0128] The sequence numbers of the above embodiments of the present invention are for descriptive purposes only and do not represent the superiority or inferiority of the embodiments.
[0129] In the above embodiments of the present invention, the descriptions of each embodiment have different focuses. For parts not described in detail in a certain embodiment, please refer to the relevant descriptions of other embodiments.
[0130] In the several embodiments provided in this application, it should be understood that the disclosed technical content can be implemented in other ways. The device embodiments described above are merely illustrative; for example, the division of units can be a logical functional division, and in actual implementation, there may be other division methods. For instance, multiple units or components may be combined or integrated into another system, or some features may be ignored or not executed. Furthermore, the displayed or discussed mutual coupling, direct coupling, or communication connection may be through some interfaces; the indirect coupling or communication connection between units or modules may be electrical or other forms.
[0131] The units described as separate components may or may not be physically separate. The components shown as units may or may not be physical units; that is, they may be located in one place or distributed across multiple units. Some or all of the units can be selected to achieve the purpose of this embodiment according to actual needs.
[0132] Furthermore, the functional units in the various embodiments of the present invention can be integrated into one processing unit, or each unit can exist physically separately, or two or more units can be integrated into one unit. The integrated unit can be implemented in hardware or as a software functional unit.
[0133] If the integrated unit is implemented as a software functional unit and sold or used as an independent product, it can be stored in a non-volatile storage medium. Based on this understanding, the technical solution of the present invention, in essence, or the part that contributes to the prior art, or all or part of the technical solution, can be embodied in the form of a software product. This computer software product is stored in a storage medium and includes several instructions to cause a computer device (which may be a personal computer, server, or network device, etc.) to execute all or part of the steps of the methods described in the various embodiments of the present invention. The aforementioned storage medium includes various media capable of storing program code, such as USB flash drives, read-only memory (ROM), random access memory (RAM), portable hard drives, magnetic disks, or optical disks.
[0134] The above description is only a preferred embodiment of the present invention. It should be noted that for those skilled in the art, several improvements and modifications can be made without departing from the principle of the present invention, and these improvements and modifications should also be considered within the scope of protection of the present invention.
Claims
1. A method for updating model parameters, characterized in that, include: Obtain the target function parameters for each of the multiple clients; Receive the representation matrix corresponding to each of the multiple clients; Based on the representation matrix corresponding to each of the multiple clients, the similarity value between any two clients among the multiple clients is calculated, and the similarity value between the multiple clients is determined. Based on the similarity values among the multiple clients and the objective function parameters corresponding to each client, the target global model parameters are determined; The target global model parameters are distributed to the multiple clients.
2. The method according to claim 1, characterized in that, Receiving the representation matrices corresponding to each of the multiple clients includes: A sample medical dataset is distributed to the multiple clients, wherein each client performs forward propagation on the sample medical dataset based on its corresponding local neural network model to obtain its own representation matrix. Receive the respective representation matrices fed back by the multiple clients.
3. The method according to claim 1, characterized in that, Based on the representation matrices corresponding to each of the plurality of clients, a similarity value is calculated between the first client and the second client among the plurality of clients, wherein the first client and the second client are any two clients among the plurality of clients, including: Based on the representation matrices corresponding to the first client and the second client respectively, determine the sample point set corresponding to the first client and the second client respectively, wherein a row vector in the representation matrix represents a sample point; Based on the sample point sets corresponding to the first client and the second client respectively, calculate the Jensen-Shannon divergence between the first client and the second client; The similarity value between the first client and the second client is determined based on the Jensen-Shannon divergence between the first client and the second client.
4. The method according to claim 1, characterized in that, The determination of the target global model parameters based on the similarity values among the multiple clients and the objective function parameters corresponding to each client includes: Based on the similarity values among the multiple clients, the aggregation weight corresponding to each of the multiple clients is determined; Based on the aggregate weights and objective function parameters corresponding to each of the multiple clients, the initial global model parameters are determined. The initial global model parameters are updated using the gradient descent algorithm to obtain the target global model parameters.
5. The method according to claim 4, characterized in that, The determination of initial global model parameters based on the aggregate weights corresponding to each of the multiple clients and the objective function parameters corresponding to each of the multiple clients includes: Clients whose aggregate weight exceeds a preset weight threshold are identified as relevant clients; The initial global model parameters are determined based on the aggregate weights corresponding to the relevant clients and the objective function parameters corresponding to the relevant clients.
6. The method according to claim 1, characterized in that, Also includes: Repeat the process of determining the target global model parameters and sending the target global model parameters to the multiple clients until the accuracy of the local neural network model corresponding to the target client in the multiple clients reaches the preset accuracy threshold. Then stop repeating and determine the current target global model parameters as the target model parameters corresponding to the target client.
7. A model parameter update device, characterized in that, include: The acquisition module is used to obtain the target function parameters corresponding to each of the multiple clients. The receiving module is used to receive the representation matrix corresponding to each of the multiple clients; The calculation module is used to calculate the similarity value between any two clients among the multiple clients based on the representation matrix corresponding to each of the multiple clients, and to determine the similarity value between the multiple clients. The determination module is used to determine the target global model parameters based on the similarity values among the multiple clients and the target function parameters corresponding to each of the multiple clients; The sending module is used to send the target global model parameters to the multiple clients.
8. A non-volatile storage medium, characterized in that, The non-volatile storage medium includes a stored program, wherein, when the program is executed, it controls the device where the non-volatile storage medium is located to execute the model parameter update method according to any one of claims 1 to 6.
9. A computer device, characterized in that, include: Memory and processor The memory stores computer programs; The processor is configured to execute a computer program stored in the memory, wherein when the computer program is executed, the processor performs the model parameter update method according to any one of claims 1 to 6.
10. A computer program product, comprising a computer program, characterized in that, When the computer program is executed by the processor, it implements the model parameter update method according to any one of claims 1 to 6.