Federated Machine Learning for Inducing Sparsity
The method induces sparsity in federated learning by using gate probability distributions to train model subsets on low-power devices, reducing communication costs and maintaining model performance, addressing the inefficiencies of existing federated learning methods.
Patent Information
- Application Number
- JP2023517950
- Authority / Receiving Office
- JP · JP
- Patent Type
- Patents
- Current Assignee / Owner
- Priority Date
- 2020-09-28
- Filing Date
- 2021-09-28
- Publication Date
- 2025-07-11
- Estimated Expiration
- 2041-09-28
AI Technical Summary
Federated learning methods face challenges in balancing model performance with communication efficiency, often resulting in sub-optimal models due to irreversible compression of data, which increases communication costs and reduces performance.
A method for federated learning that induces sparsity by generating model elements based on gate probability distributions, allowing clients to train subsets of the global model, and communicating only the updated values and gate probabilities, optimizing the global model through client-specific probabilities.
This approach reduces communication costs while maintaining model performance by enabling efficient training of sparse models, facilitating training on low-power devices and minimizing data exchange without sacrificing accuracy.
Smart Images

Figure 0007706544000045 
Figure 0007706544000046 
Figure 0007706544000047
Abstract
Description
Technical Field
[0001]
[0001] Cross - reference to Related Applications This application claims the benefit and priority of Greek Patent Application No. 20200100587, filed on September 28, 2020, the entire content of which is incorporated herein by reference.
[0002]
[0002] Aspects of the present disclosure relate to sparsity - inducing federated machine learning.
Background Art
[0003]
[0003] Machine learning is generally a process of generating a trained model (e.g., an artificial neural network, a tree, or other structure), which represents a generalized fit to a set of training data. Applying the trained model to new data generates inferences that can be used to gain insights into the new data.
[0004]
[0004] In various technical fields related to what is sometimes called artificial intelligence tasks, as the use of machine learning has been increasing rapidly, there has been a need for more efficient processing of machine learning model data. For example, "edge processing" devices such as mobile devices, always - on devices, and Internet of Things (IoT) devices need to balance the implementation of advanced machine learning capabilities with various interrelated design constraints such as packaging size, native computing power, power storage and usage, data communication capabilities and costs, memory size, and heat dissipation.
[0005]
[0005] Federated learning is a distributed machine learning framework that enables several clients, such as edge processing devices, to collaboratively train a shared global model without transferring their local data to a remote server. Generally, a central server coordinates the federated learning process, and each participating client communicates only model parameter information with the central server while keeping its local data private. This distributed approach helps with the issue of client device capacity limitations (since training is federated) and also, in many cases, alleviates concerns about data privacy.
[0006]
[0006] Federated learning generally limits the amount of model data in any single transmission between a server and a client (or vice versa), but the iterative nature of federated learning still generates a significant amount of data transmission traffic during training, which can be extremely costly depending on the device and connection type. Therefore, it is generally desirable to reduce the size of data exchange between the server and the client during federated learning. However, conventional methods for reducing data exchange have resulted in less performant models, such as when irreversible compression of model data is used to limit the amount of data exchanged between the server and the client.
[0007]
[0007] Therefore, there is a need for an improved method of performing federated learning in which model performance is not sacrificed for communication efficiency.
Summary of the Invention
[0008]
[0008] Some aspects provide a method for performing federated learning of a machine learning model, which includes, for each of a plurality of clients and for each of a plurality of training rounds, generating, for each model element of a set of model elements for a global machine learning model, a subset of model elements for each client based on sampling a gate probability distribution for the model element; sending to each client the subset of model elements and a set of gate probabilities based on the sampling, wherein each gate probability of the set of gate probabilities is associated with one of the model elements of the subset of model elements; receiving, from each of the plurality of clients, a respective set of model update values; and updating the global machine learning model based on the respective sets of model update values received from each of the plurality of clients.
[0009]
[0009] Further aspects provide a method for performing federated learning of a machine learning model, which includes receiving, from a server that manages federated learning of a global machine learning model, a subset of model elements from a set of model elements for the global machine learning model and a set of gate probabilities, wherein each gate probability of the set of gate probabilities is associated with one of the model elements of the subset of model elements; generating a set of model update values based on training a local machine learning model based on the set of model elements and the set of gate probabilities; and sending the set of model update values to the server.
[0010]
[0010] Another aspect provides a processing system configured to execute the above-described method and the methods described herein, a non-transitory computer-readable medium comprising instructions that, when executed by one or more processors of the processing system, cause the processing system to execute the above-described method and the methods described herein, a computer program product embodied on a computer-readable storage medium comprising code for executing the above-described method and the methods further described herein, and a processing system comprising means for executing the above-described method and the methods further described herein.
[0011]
[0011] The following description and the related drawings detail some exemplary features of one or more embodiments.
[0012]
[0012] The accompanying drawings illustrate some aspects of one or more embodiments and, accordingly, should not be regarded as limiting the scope of the disclosure.
Brief Description of the Drawings
[0013]
Figure 1
[0013] A diagram illustrating an exemplary training flow for promoting sparsity in collaborative learning.
Figure 2
[0014] A diagram illustrating an exemplary method for performing collaborative learning that induces sparsity.
Figure 3
[0015] A diagram illustrating another exemplary method for performing collaborative learning that induces sparsity.
Figure 4
[0016] A diagram illustrating an exemplary processing system that may be configured to execute aspects of the collaborative learning methods described herein.
Modes for Carrying Out the Invention
[0014]
[0017] For ease of understanding, where possible, the same reference numbers are used to denote identical elements common to the drawings. It is contemplated that the elements and features of one embodiment may be beneficially incorporated into other embodiments without further recitation.
[0015]
[0018] Aspects of the present disclosure provide an apparatus, method, processing system, and computer-readable medium for federated machine learning that induces sparsity.
[0016]
[0019] As machine learning models become more complex and thus larger, it is increasingly difficult to train them on anything other than high-performance computers such as servers. Federated learning is a distributed machine learning framework that enables several clients, including low-power devices such as edge processing devices, to collaboratively train a shared global model. In such a setting, it is generally desirable to reduce client device computations along with the total communication cost. In particular, high communication costs can render federated learning over mobile data impractical.
[0017]
[0020] One approach to address these challenges is "federated dropout" where the server selects a specific probability of selecting submodels from the original model prior to the federated training process. Then, during the training process, the server probabilistically selects random submodels and communicates them to each client. Thus, instead of training update values for the entire global model locally, each client trains update values for a smaller submodel. Since the submodels are subsets of the global model, the local update values computed by the clients have a natural interpretation as update values for the larger global model.
[0018]
[0021] Another approach is to modify the messages from the client to the server for the economy of data communication. For example, the client may select the top k most beneficial elements from the message for the server and communicate only those k most beneficial elements to the server. Alternatively, the client may quantize the message before it is communicated to the server.
[0019]
[0022] The embodiments described herein improve upon existing techniques in several important ways. First, unlike conventional federated dropout techniques, the methods described herein enable each client to automatically determine an appropriate submodel of the original model in a way that is as efficient as possible while fitting its local dataset. Second, rather than the server adhering to one particular global probability across submodels, the global model can be optimized through client-specific probabilities.
[0020] Federated Averaging Through the Lens of Expectation Maximization
[0023] As described above, federated learning generally operates on a dataset D = {(x1, y1), …, (x N , y N )} of N data points potentially distributed non-iid across S shards without direct access to shard-specific datasets, i.e., D = D1 ∪ … ∪ D STo address the problem of learning a server model (e.g., a neural network) using parameter w, where parameter w can generally represent a vector, matrix, or tensor. A shard may generally be a client for processing participating in federated learning with a central server, and it should be noted that a shard may comprise a remote computer, server, mobile device, smart device, edge processing device, etc. Without loss of generality and for simplicity, it is assumed hereinafter that all of shard S has the same amount of data points, although the framework can be extended to non-uniform amounts of data points by selecting appropriate weighting factors. The loss function L s (D s ;w) is defined, and the total loss can be described as follows.
[0021]
Number
[0022]
[0024] Here, N s is the number of data points in shard (e.g., device) s, and D s is the data set in the device of shard s. In particular, this objective corresponds to minimizing the empirical risk (ERM: empirical risk minimization) over the joint data set D with loss L(·) for each data point.
[0023]
[0025] It is desirable to reduce the communication cost of federated learning. One approach to reducing communication during federated learning is to perform multiple gradient updates with respect to w for the purpose of internal optimization for each shard s, and thus obtain a "local" model with parameter φ s These multiple gradient update values are shown as the number of passes through the entire local data set, abbreviated as E, i.e., "local epochs". Then each of the shards has a local (or sub) model φ sCommunicate it to the server, and the server updates the global model at "round" t, for example, by averaging the parameters of the local machine learning model according to the following formula.
[0024]
Number
[0025]
[0026] This approach is sometimes called federated averaging.
[0026]
[0027] Although it is easy to implement, federated averaging can provide sub-optimal results for non-IID data, even if its convergence can be proven. In fact, when the shard S has a skewed distribution, the average of the local machine learning model parameters can sometimes be a poor estimate of the global model. To counter this, a "proximal" term for optimization at the shard level may be used, and the proximal term promotes the local machine learning model φ s to "approach" the model w at the server under a certain distance. More formally, this can be defined as follows.
[0027]
Number
[0028]
[0028] Here,
[0029]
Number
[0030] is the proximal term. After each shard-specific optimization is completed, the global model can be updated in a similar way to federated averaging, that is, by averaging the shard-specific parameters in Equation 2.
[0031] Connecting Federated Averaging with Expectation Maximization
[0029] In particular, the entire federated averaging algorithm is compatible with an optimization procedure based on a given objective function. For example, consider the following objective function.
[0032]
Equation
[0033]
[0030] Here, D s corresponds to a shard-specific dataset with N s data points, and p(D s | w) corresponds to the likelihood of D s under the server parameters w and Σ s N s = N. Next, consider decomposing each of the shard-specific likelihoods as follows.
[0034]
Equation
[0035]
[0031] Here, an auxiliary latent variable φ s is introduced, and the server parameter w acts as a hyperparameter of the prior distribution over the shard-specific parameter p(φ s | w). These latent variables are the parameters of the local machine learning model in shard s, and the following convenient form of the prior distribution can be used.
[0036]
Equation
[0037]
[0032] Here, λ acts as a regularization strength to prevent φ s from being too far from w. Then, overall, this leads to the following objective function.
[0038]
Number
[0039]
[0033] Latent variable φ s One way to optimize this objective in the presence of is by Expectation Maximization (EM). EM generally consists of two steps. An expectation step in which the posterior distribution is formed over the latent variables, i.e.,
[0040]
Number
[0041]
[0034] And, with respect to the model parameters w, by marginalizing over this posterior distribution as follows, D s This is the maximization step in which the probability of is maximized.
[0042]
Number
[0043]
[0035] Thus, if a single gradient step is performed on w in the maximization step, this procedure corresponds to performing gradient descent on the original objective of Equation 7. To illustrate this, the gradient of Equation 7 can be taken with respect to w, which is s Z = ∫ p(D s | φ s ) p(φ s | w) dφ s as given by
[0044]
Number
[0045]
[0036] Here, to compute Equation 12, the local variable φ sThe posterior distribution of [[]] must first be obtained, and then the gradient of w is estimated by marginalizing over this posterior distribution.
[0046]
[0037] When posterior inference is difficult to solve, hard EM is sometimes employed. In such cases, for the latent variable φ s the "hard" assignment of [[]] can be done, in the expectation step, for example, by approximating p(φ s |D s ) at its most likely point.
[0047]
Number
[0048]
[0038] This is usually easier to do using techniques such as stochastic gradient ascent. Given these hard assignments, the maximization step corresponds to another simple maximization of the following equation.
[0049]
Number
[0050]
[0039] As a result, hard EM corresponds to a block coordinate ascent type algorithm for the following objective function.
[0051]
Number
[0052]
[0040] Here, optimizing φ 1:S while fixing w is done alternately with optimizing w while fixing φ 1:S .
[0053]
[0041] By setting λ → 0 in Equation 6, it is clear that the hard assignment in the expectation step mimics the process of optimizing the local machine learning model on each shard. In fact, even by locally optimizing the model using stochastic gradient descent for a fixed number of iterations with a given learning rate, a specific prior distribution can be assumed over the parameters. In the case of linear regression, this prior distribution is a Gaussian distribution centered around the initial values of the parameters, and in the case of non - linear models, this prior distribution can be shown through the proximal view of each gradient descent iteration.
[0054]
Number
[0055]
[0042] This imposes a similar prior Gaussian distribution centered around the previous iteration, and the learning rate η acts as the variance of that prior distribution. After obtaining φ * s then the maximization step corresponds to the following equation.
[0056]
Number
[0057]
[0043] Then, the closed - form solution for this objective can be found by setting the derivative of the objective with respect to w to zero and solving for w according to the following equation.
[0058]
Number
[0059]
[0044] Here, the optimal solution for w when φ * 1:S is given is the same as the average of φ generated using associative averaging. * 1:S
[0060]
[0045] Federated averaging does not optimize the local parameter φ in each round. However, the alternating procedure of EM corresponds to block coordinate ascent on a single objective function that is a variational lower bound of the log marginal likelihood. More specifically, the EM iteration performs block coordinate ascent to optimize the following objective. s While federated averaging does not optimize the local parameter φ towards convergence in each round, the alternating procedure of EM corresponds to block coordinate ascent on a single objective function that is a variational lower bound of the log marginal likelihood. More specifically, the EM iteration performs block coordinate ascent to optimize the objective given by the following equation.
[0061]
Number
[0062]
[0046] Here, w s is the parameter of the variational approximation to the posterior distribution p(φ s | D s , w). To obtain the procedure of federated averaging up to machine precision, a deterministic distribution of φ s , i.e.,
[0063]
Number
[0064] may be used, which leads to a simplification of the following objective.
[0065]
Number
[0066]
[0047] Here, C is a fixed constant independent of the parameter to be optimized. In particular, this objective is the same as the objective in Equation 15.
[0067] Encouraging Sparsity in Federated Learning
[0048] Strengthening the joint averaging promotes sparsity through an appropriate prior distribution. Promoting sparsity has two important advantages. First, the model becomes smaller and thus it is easier to train on the device in terms of hardware. Second, since pruned parameters do not need to be communicated, the communication cost is reduced.
[0068]
[0049] The standard for sparsity in a Bayesian model is the spike-slab prior distribution. This is a mixture of two components: a delta spike δ(0) at zero and a continuous distribution over the real line, i.e., a slab. More specifically, the spike-slab prior distribution can be defined as follows for a Gaussian slab.
[0069]
Number
[0070]
[0050] Or it can be equivalently defined as a hierarchical model as follows.
[0071]
Number
[0072]
[0051] Here, z serves as a "gating" variable that switches the parameter w on or off. Next, consider using this distribution instead of a single Gaussian distribution for the prior distribution over the parameters in the joint setting. In this case, the hierarchical model becomes as follows.
[0073]
Number
[0074]
[0052] Here, w is the model weight in the server, and θ is the probability of the binary gate. Similar to the joint averaging, hard EM can be performed to optimize w and θ using the approximate distribution q(φ s |z s )q(z s ). Then, the variational lower bound of this model can be described as follows.
[0075]
Number
[0076]
[0053] Or, equivalently, it can be described as the following equation.
[0077]
Number
[0078]
[0054] Regarding the shard specific weight distribution, since they are continuous, q(φ si |z si = 1):=N(φ si ,ε), q(φ si |z si = 1):=N(0,ε) may be used with ε≒0, and the shard specific weight distribution is deterministic up to machine precision. However, for the gating variable, since it is binary,
[0079]
Number
[0080] is used together with π si which is the probability of activating the local gate z si , where Bern(·) represents the Bernoulli distribution. To perform hard EM for binary variables,
[0081]
Number
[0082] The entropy term of can be removed from the above boundary because the approximate distribution is promoted to move towards the most probable value of z s Furthermore, to reach a simple and intuitive objective at the shard level, the spike at zero can be relaxed to a Gaussian distribution with precision λ2, i.e., p(φ si |z si =0)=N(0,1 / λ2). Taking all these into account, by connecting appropriate equations to Equation 26, it can be shown that the local and global objectives become the following equations respectively.
[0083]
Number
[0084]
[0055] Here,
[0085]
Number
[0086] and C are constants independent of the variables to be optimized. In particular, locally, each shard optimizes its weights to be close to the server weights adjusted by the prior precision λ and its probability π s while explaining D as much as possible. Furthermore, the gate activation probability is optimized to be close to the server θ using an additional term that penalizes the sum of the local activation probabilities. This is the same as the previously proposed L0 regularization objective. s
[0087]
[0056] Next, the local shard passes through some procedure for φ s and π s After optimization, what happens on the server can be considered. Since the server loss with respect to \(w\) and \(\theta\) is just the sum of all local losses, the gradient for each of the parameters is as follows.
[0088]
Number
[0089]
[0057] Setting these derivatives to zero, the stationary points are as follows.
[0090]
Number
[0091]
[0058] That is, it is the weighted average of the local weights and the average of the local probabilities that hold these weights. Therefore, \(\pi\) s is optimized to be sparse through the \(L_0\) penalty, so the server probability \(\theta\) will also be sparse for weights that are not used by any of the shards. As a result, to obtain the final sparse architecture, the weights can be pruned when their server inclusion probability \(\theta\) is less than a threshold such as 0.1, although other thresholds are possible.
[0092] Local Optimization
[0059] Optimizing \(\varphi_s\) locally is straightforward using a gradient-based optimizer, but the expected value with respect to the binary variable \(z\) in Equation 27 s is difficult to compute in closed form, and using Monte Carlo integration does not yield reparameterizable samples, so \(\pi\) s is not straightforward. To avoid these issues, the objective can be rewritten in an equivalent form as follows.
[0093]
Number
[0094]
[0060] Next, it may be replaced with a continuous relaxation such as a hard-Concrete distribution. The continuous relaxation
[0095]
Math
[0096] is done as follows, where v
[0097]
Math
[0098] is a parameter of the surrogate distribution. In this case, the local objective becomes as follows. s Here,
[0099]
Math
[0100]
[0061] Here,
[0101]
Math
[0102] is the continuous relaxation,
[0103]
Math
[0104] is the cumulative distribution function (CDF) of. Therefore, next, the surrogate objective can be easily optimized using gradient descent.
[0105] Reducing the Client to Server Communication Cost
[0062] The above model enables learning a sparse model for inference at the server. The same framework can be used to reduce the communication cost during training by adopting two techniques that reduce the communication cost for each of the communication from the client to the server and from the server to the client, respectively.
[0106]
[0063] To reduce the cost from the client to the server, instead of the distribution itself, sparse samples from the local distribution can be communicated. For example, instead of sending the local weights φ s and the local probability π s to the server, the client can, instead, draw a random binary sample z s ∈ {0, 1} according to π s , and then communicate only the weights φ si that have z si = 1 to the server together with z s . In this way, the zero values of the parameter vector do not need to be communicated, which results in a significant saving while still keeping the server gradient unbiased. More specifically, the gradient and the stationary point of the server weights can be expressed as follows.
[0107]
Equation
[0108]
[0064] On the other hand, for the equation of the server probability, it is as follows.
[0109]
Equation
[0110]
[0065] As a result, the client can send a subset of the local weights
[0111]
Number
[0112] only
[0113]
Number
[0114] can communicate via. In this way, the client can communicate with a subset of the local weights z s along with. When accessing these samples, the client can form a one-sample probability estimate of either the gradient or the stationary point of w, θ. The client locally operates for the purpose of smoothing using hard-Concrete relaxation
[0115]
Number
[0116] since it operates for the purpose of smoothing using
[0117]
Number
[0118] is, when the client communicates with the server, thus, whenever obtaining an exact discrete sample z s can be formed by sampling from zero temperature
[0119]
Number
[0120] from.
[0121] Note that this is a method of reducing communication volume without adding bias to the gradient of the original objective. If it is acceptable to receive additional bias, further techniques such as quantization and top-k gradient selection can be used to further reduce communication volume.
[0122] Reducing the Server to Client Communication Cost
[0067] The server needs to communicate the updated distribution to the client in each round. Unfortunately, in the case of simple unstructured pruning, for each weight w i there is a related θ i that needs to be sent to the client, which doubles the communication cost. To mitigate this effect, a single additional parameter indicating the probability for each group of weights is introduced, and thus structured pruning, which is more efficient in terms of the number of trainable parameters compared to unstructured pruning, can be adopted. Even with structured pruning, the normal weights and probabilities are sent to the server (as mentioned above, when communicating sparse samples, the probability vector becomes extremely small when using structured pruning). Therefore, for a moderately sized group, for example, the set of weights of a given convolution filter, the extra overhead is relatively small.
[0123]
[0068] If some bias is allowed in the optimization procedure, the reduction of communication cost can be further advanced. For example, the global model is pruned during training after each round, and thus only a subset of the surviving models can be sent to each of the clients. In particular, this is efficient to execute and does not require any data at the server, because the server has access to the inclusion probability θ, and thus parameters with θ less than a threshold, for example, less than 0.1, can be removed. This can result in a substantial reduction of communication cost, especially during the later stages of training when the model is sparser.
[0124]
[0069] An additional way to reduce communication cost is for the client to perform local pruning and thus request from the server only a subset of the original model parameters that will survive locally.
[0125]
[0070] Thus, when performing federated learning, a generalization of federated averaging may be used to optimize sparse neural networks, and the generalization of federated averaging subsequently leads to significant communication savings while maintaining similar performance.
[0126] Example Training Flow for Encouraging Sparsity in Federated Learning
[0071] FIG. 1 shows an exemplary training flow for encouraging sparsity in federated learning, as conceptually described in detail above.
[0127]
[0072] First, server 102 generates or maintains the global model 104 in a first state. In this example, each edge between nodes in the global model 104 is associated with a parameter (e.g., parameter set 105) that includes a weight w and a gate probability θ. As described above, the gate probability generally represents the likelihood that the associated weight is included in a local (or sub) model for federated training.
[0128]
[0073] At 110, server 102 samples the global model weights w according to their associated gate probabilities θ to generate various subsets of weights and gate probabilities for each of shards 106A - K, where each shard may represent a client device participating in federated learning with server 102.
[0129]
[0074] Based on this information, each of shards 106A - K, where K is the total number of shards participating in federated learning, generates local machine learning models 108A - K using parameters φ s , π s based on the parameters received from server 102, where s is a particular shard within the set S of shards. In FIG. 1, the dotted lines between nodes within local machine learning models 108A - K are gated off and thus indicate weights not included in local machine learning model training.
[0130]
[0075] As shown, local machine learning models generally vary from shard to shard based on different gate probabilities and random sampling performed by server 102. This helps increase the inclusivity of federated training.
[0131]
[0076] In 112, each of the shards 106A - K trains its respective local machine learning model 108A - K, generating updated local machine learning models 108A' - K'. Further, each of the shards 106A - K generates a weight gradient and a gate gradient based on the training, for example, as described above with respect to equations 31 and 32.
[0132]
[0077] In 114, each of the shards 106A - K returns model update data to the server 102. The server 102 then uses the model update data to generate an updated global model 104'. In the illustrated embodiment, the model update data sent by each of the shards 106A - K includes the weight gradients and gate gradients for each element of the shard's local machine learning model (e.g., 108A' - K').
[0133]
[0078] In particular, FIG. 1 shows a single round of training for simplicity, and this process can be repeated any number of times, for example, until a training target is reached (e.g., the number of iterations is completed, the weights converge, an accuracy threshold is reached, etc.).
[0134]
[0079] After the collaborative training is completed (e.g., when the global model 104 converges), one or more nodes (in the example of a neural network model) can be permanently and effectively gated off (not shown in FIG. 1). More generally, the pruning rate of the global model 104 can be gradually increased during training so that the model can become extremely sparse (e.g., a sparsity rate of about 90%) by the end of training. For example, a 90% sparsity rate of the trained global model 104' in the context of FIG. 1 means that 90% of the weights are pruned during training based on a set threshold.
[0135]
[0080] In particular, in this example, sparsity is induced in the weights for the edges between the nodes of the exemplary model, but in other examples, other aspects of the model can be associated with the gate probabilities to induce alternative or additional sparsity. For example, nodes or layers within the model can be associated with gate probabilities and thus can be sampled and pruned during collaborative training. As another example, in the context of a convolutional neural network model, individual filter channels can be associated with gate probabilities and thus can be sampled and pruned to induce sparsity during training.
[0136]
[0081] In addition to the sparsity induced during training based on the gate probabilities, additional strategies can be implemented to reduce communication costs. As described above, in order to reduce the communication cost from the shard (or client) to the server (e.g., in step 114), only the gradients for the aspects of the model that are not gated off (e.g., the weights represented by the solid lines between the nodes in FIG. 1) are sent back to the server during each training round. Thus, unlike conventional collaborative learning where all weights are transmitted between the shard and the server in each training round, here it is possible to save communication time and cost by sending only a subset of the model data corresponding to what is updated by each local machine learning model 108A - K during local training.
[0137]
[0082] Further, each shard (e.g., 106A - K) can sample elements of the local machine learning model (e.g., 108A - K) according to the gate probability π s Accordingly, for example, the weight gradients (for the parameters φ s of the local machine learning model) and (the local gate probability π sRather than sending the entire set of gate gradients for), the shard can either send the weight update value and z = 1, or send nothing (corresponding to z = 0), where, as described above, z is the "gating" variable. Thus, z is a value in {0, 1}, π is the probability of having z = 1, and 1 - π is the probability of having z = 0.
[0138]
[0083] This helps reduce the communication cost between each shard and the server 102 in step 114. In such a case, the server update rule can be modified from equations (30) to equations (34) and (36) to update the weight w and the probability for the binary gate, respectively.
[0139] Example Methods of Performing Federated Learning
[0084] Figure 2 shows an exemplary method 200 for performing sparsity - inducing federated learning, which can be executed by a federated learning server such as 102 in FIG. 1, for example.
[0140]
[0085] Method 200 starts at step 202 by generating a subset of model elements for each of a plurality of clients (e.g., shards 106A - K in FIG. 1) based on sampling a gate probability distribution for each model element of a set of model elements of a global machine learning model.
[0141]
[0086] In some embodiments of method 200, a subset of model elements comprises a subset of weights associated with edges connecting nodes in the global machine learning model. In some embodiments of method 200, a subset of model elements comprises a subset of nodes in the global machine learning model. In some embodiments of method 200, a subset of model elements comprises a subset of channels in the convolutional filters of the global machine learning model.
[0142]
[0087] Method 200 then proceeds to step 204 of sending, to each respective client of the plurality of clients, a subset of model elements and a set of gate probabilities based on sampling, where each gate probability of the set of gate probabilities is associated with one of the model elements of the subset of model elements (e.g., as described in step 110 with respect to FIG. 1).
[0143]
[0088] Method 200 then proceeds to step 206 of receiving, from each respective client of the plurality of clients, a respective set of model update values (e.g., as described in step 114 with respect to FIG. 1).
[0144]
[0089] Method 200 then proceeds to step 208 of updating the global machine learning model based on the respective sets of model update values from each respective client of the plurality of clients.
[0145]
[0090] In some embodiments of method 200, each respective set of model update values comprises a set of weight gradients associated with a local machine learning model trained by each respective client and a set of gate probability gradients associated with a local machine learning model trained by each respective client.
[0146] In some embodiments of method 200, each set of model update values comprises a set of weight gradients associated with a local machine learning model trained by each client, and a binary gate variable value associated with each weight gradient of the set of weight gradients.
[0147] In some embodiments of method 200, updating the global machine learning model based on each set of model update values from each respective client of a plurality of clients further comprises pruning the updated global machine learning model based on an updated gate probability for the global machine learning model and a threshold gate probability value.
[0148] In particular, FIG. 2 is merely an example of a model consistent with the disclosure herein, and additional steps, fewer steps, and / or further examples with additional steps are possible.
[0149] FIG. 3 shows another exemplary method 300 for performing collaborative learning that induces sparsity, which may be performed by collaborative learning clients such as 106A-K of FIG. 1, for example.
[0150] Method 300 begins at step 302 of receiving, from a server that manages collaborative learning of a global machine learning model, a subset of model elements from a set of model elements for the global machine learning model and a set of gate probabilities, where each gate probability of the set of gate probabilities is associated with one model element of the subset of model elements.
[0151]
[0096] In some embodiments of method 300, a subset of model elements comprises a subset of weights associated with edges connecting nodes in the global machine learning model. In some embodiments of method 300, a subset of model elements comprises a subset of nodes in the global machine learning model. In some embodiments of method 300, a subset of model elements comprises a subset of channels in the convolutional filters of the global machine learning model.
[0152]
[0097] Method 300 then proceeds to step 304 (e.g., as described in step 112 with respect to FIG. 1) of generating a set of model update values based on training a local machine learning model based on the set of model elements and the set of gate probabilities.
[0153]
[0098] Method 300 then proceeds to step 306 (e.g., as described in step 114 with respect to FIG. 1) of sending the set of model update values to the server.
[0154]
[0099] In some embodiments of method 300, the set of model update values comprises a set of weight gradients associated with the local machine learning model and a set of gate probability gradients associated with the local machine learning model (e.g., local machine learning models 108A - K of FIG. 1).
[0155]
[0100] In some embodiments of method 300, the set of model update values comprises a set of weight gradients associated with the local machine learning model and binary gate variable values associated with each weight gradient of the set of weight gradients.
[0156]
[0101] In some embodiments, method 300 further comprises receiving a final set of model elements from the server, where the final set of model elements corresponds to a pruned global machine learning model.
[0157]
[0102] In particular, FIG. 3 is merely an example of a model that is consistent with the disclosure herein, and additional examples with additional steps, fewer steps, and / or additional steps are possible.
[0158] Exemplary Processing System
[0103] FIG. 4 shows an exemplary processing system 400 that may be configured to execute aspects of the federated learning method described herein, including, for example, methods 200 and 300 of FIGS. 2 and 3, respectively.
[0159]
[0104] The processing system 400 includes a central processing unit (CPU) 402 that may be a multi-core CPU in some examples. Instructions executed in the CPU 402 may be loaded, for example, from a program memory associated with the CPU 402 or from the memory 424.
[0160]
[0105] The processing system 400 also includes additional processing components adapted for specific functions, such as a graphics processing unit (GPU) 404, a digital signal processor (DSP) 406, a neural processing unit (NPU) 408, a multimedia processing unit 410, and a wireless connectivity component 412.
[0161]
[0106] An NPU, such as 408, is generally a dedicated circuit configured to implement control and arithmetic logic for executing machine learning algorithms, such as algorithms for processing artificial neural networks (ANNs), deep neural networks (DNNs), random forests (RFs), etc. The NPU is alternatively sometimes referred to as a neural signal processor (NSP), a tensor processing unit (TPU), a neural network processor (NNP), an intelligence processing unit (IPU), or a vision processing unit (VPU).
[0162]
[0107] NPUs such as 408 can be configured to accelerate the performance of general machine learning tasks such as image classification, voice classification, and various other prediction models. In some examples, multiple NPUs can be instantiated on a single chip such as a system-on-chip (SoC), while in other examples, these can be part of a dedicated neural network accelerator.
[0163]
[0108] The NPU can be optimized for training or inference, or in some cases, configured to balance performance between the two. In the case of an NPU capable of performing both training and inference, the two tasks can still generally be executed independently.
[0164]
[0109] An NPU designed to accelerate training is generally configured to accelerate the optimization of new models. The NPU takes as input an existing dataset (often labeled or tagged) and iterates over the dataset to adjust model parameters such as weights and biases in an operation that is highly computationally intensive in order to improve model performance. Generally, optimizing based on incorrect predictions involves backpropagating through the layers of the model and determining gradients to reduce the prediction error.
[0165]
[0110] An NPU designed to accelerate inference is generally configured to operate on a complete model. Thus, such an NPU can be configured to quickly process new data through a pre-trained model to generate a model output (e.g., an inference).
[0166]
[0111] In one implementation, the NPU 408 is part of one or more of the CPU 402, GPU 404, and / or DSP 406.
[0167]
[0112] In some examples, the wireless connectivity component 412 may include sub-components for, for example, third generation (3G) connectivity, fourth generation (4G) connectivity (e.g., 4G LTE (registered trademark)), fifth generation connectivity (e.g., 5G or NR), Wi-Fi (registered trademark) connectivity, Bluetooth (registered trademark) connectivity, and other wireless data transmission standards. The wireless connectivity processing component 412 is further connected to one or more antennas 414.
[0168]
[0113] The processing system 400 may also include one or more sensor processing units 416 associated with sensors of any type, one or more image signal processors (ISPs) 418 associated with image sensors of any type, and / or a navigation processor 420 that may include satellite-based positioning system components (e.g., GPS or GLONASS) and inertial positioning system components.
[0169]
[0114] The processing system 400 may also include one or more input and / or output devices 422 such as a screen, a touch sensor surface (including a touch sensor display), physical buttons, speakers, microphones, etc.
[0170]
[0115] In some examples, one or more of the processors of the processing system 400 may be based on the ARM or RISC-V instruction set.
[0171]
[0116] The processing system 400 also includes a memory 424 representing one or more static memories and / or dynamic memories such as dynamic random access memory, flash-based static memory, etc. In this example, the memory 424 includes computer-executable components that may be executed by one or more of the above processors of the processing system 400.
[0172]
[0117] In this example, the memory 424 includes a transmission component 424A, a reception component 424B, a training component 424C, an inference component 424D, a sampling component 424E, a pruning component 424F, model parameters 424G (e.g., the weights and gate probabilities described above), and a model 424H. The illustrated components, as well as other components not shown, may be configured to perform various aspects of the methods described herein.
[0173]
[0118] The processing system 400 is merely an example and generally may perform the operations of the servers and / or clients / shards described herein. However, in other embodiments, some aspects may be omitted. For example, the server may omit some features that may typically be found in a mobile device, such as the multimedia component 410, the wireless connectivity component 412, the antenna 414, the sensor 416, the ISP 418, and the navigation component 420. The illustrated example is not meant to be limiting.
[0174] Exemplary Clauses
[0119] Implementation examples are described in the following numbered clauses.
[0175]
[0120] Clause 1: A method for performing collaborative learning of a machine learning model, comprising, for each respective client of a plurality of clients and for each training round of a plurality of training rounds, generating, based on sampling a gate probability distribution for each model element of a set of model elements for a global machine learning model, a subset of model elements for each respective client; sending to each respective client a subset of model elements and a set of gate probabilities based on the sampling, wherein each gate probability of the set of gate probabilities is associated with one model element of the subset of model elements; receiving, from each respective client of the plurality of clients, a respective set of model update values; and updating the global machine learning model based on the respective sets of model update values from each respective client of the plurality of clients.
[0176]
[0121] Clause 2: The method of clause 1, wherein the subset of model elements comprises a subset of weights associated with edges connecting nodes in the global machine learning model.
[0177]
[0122] Clause 3: The method of clause 2, wherein each respective set of model update values comprises a set of weight gradients associated with a local machine learning model trained by each respective client and a set of gate probability gradients associated with the local machine learning model trained by each respective client.
[0178]
[0123] Clause 4: The method of clause 2, wherein each respective set of model update values comprises a set of weight gradients associated with a local machine learning model trained by each respective client and a binary gate variable value associated with each weight gradient of the set of weight gradients.
[0179]
[0124] Clause 5: The method according to any one of clauses 1 to 4, wherein the subset of model elements comprises a subset of nodes in the global machine learning model.
[0180]
[0125] Clause 6: The method according to any one of Clauses 1 to 5, wherein a subset of model elements comprises a subset of channels within a convolutional filter of a global machine learning model.
[0181]
[0126] Clause 7: The method according to any one of Clauses 1 to 6, wherein updating the global machine learning model based on each respective set of model update values from a plurality of clients further comprises pruning the updated global machine learning model based on an updated gate probability for the global machine learning model and a threshold gate probability value.
[0182]
[0127] Clause 8: A method for performing federated learning of a machine learning model, comprising receiving, from a server managing federated learning of a global machine learning model, a subset of model elements from a set of model elements for the global machine learning model and a set of gate probabilities, wherein each gate probability of the set of gate probabilities is associated with one model element of the subset of model elements; generating a set of model update values based on training a local machine learning model based on the set of model elements and the set of gate probabilities; and transmitting the set of model update values to the server.
[0183]
[0128] Clause 9: The method according to Clause 8, wherein the subset of model elements comprises a subset of weights associated with edges connecting nodes within the global machine learning model.
[0184]
[0129] Clause 10: The method according to Clause 9, wherein the set of model update values comprises a set of weight gradients associated with the local machine learning model and a set of gate probability gradients associated with the local machine learning model.
[0185] Clause 11: The set of model update values is the method described in Clause 9, comprising a set of weight gradients associated with a local machine learning model and binary gate variable values associated with each weight gradient in the set of weight gradients.
[0186]
[0131] Clause 12: The subset of model elements is the method described in any one of Clauses 8 to 11, comprising a subset of nodes within a global machine learning model.
[0187]
[0132] Clause 13: The subset of model elements is the method described in any one of Clauses 8 to 11, comprising a subset of channels within the convolutional filter of a global machine learning model.
[0188]
[0133] Clause 14: Further comprising receiving a final set of model elements from a server, the final set of model elements corresponding to a pruned global machine learning model, the method described in any one of Clauses 8 to 13.
[0189]
[0134] Clause 15: A processing system comprising a memory having computer-executable instructions and one or more processors, the one or more processors being configured to execute the computer-executable instructions to cause the processing system to execute the method described in any one of Clauses 1 to 14.
[0190]
[0135] Clause 16: A processing system comprising means for executing the method described in any one of Clauses 1 to 14.
[0191]
[0136] Clause 17: A non-transitory computer-readable medium having computer-executable instructions, which, when executed by one or more processors of a processing system, cause the processing system to execute the method described in any one of Clauses 1 to 14.
[0192]
[0137] A computer program product embodied on a computer-readable storage medium comprising code for performing the method according to any one of clauses 1 to 14.
[0193] Additional Considerations
[0138] The above description has been provided to enable a person skilled in the art to make and use various embodiments described herein. The examples described herein are not intended to limit the scope, applicability, or embodiments described in the claims. Various modifications to these embodiments will be readily apparent to those skilled in the art, and the general principles defined herein may be applied to other embodiments. For example, changes may be made in the function and configuration of the elements described without departing from the scope of the present disclosure. Various examples may, as appropriate, omit, substitute, or add various procedures or components. For example, the methods described may be performed in an order different from that described, and various steps may be added, omitted, or combined. Also, the features described with respect to some examples may be combined in some other examples. For example, the apparatus may be implemented or the method may be performed using any number of the aspects described herein. Further, the scope of the present disclosure is intended to cover such apparatus or methods implemented using other structures, functions, or structures and functions in addition to, or other than, the various aspects of the present disclosure described herein. It should be understood that any aspect of the present disclosure disclosed herein may be implemented by one or more elements of the claims.
[0194]
[0139] As used herein, the term "exemplary" means "serving as an example, instance, or illustration." Any aspect described herein as "exemplary" should not necessarily be construed as preferred or advantageous over other aspects.
[0195] As used herein, the phrase "at least one of" in a list of items refers to any combination of those items including a single member. By way of example, "at least one of a, b, or c" includes a, b, c, a-b, a-c, b-c, and a-b-c, as well as any combination having multiple of the same element (e.g., a-a, a-a-a, a-a-b, a-a-c, a-b-b, a-c-c, b-b, b-b-b, b-b-c, c-c, and c-c-c, or any other order of a, b, and c).
[0196] As used herein, the term "determining" encompasses a wide variety of actions. For example, "determining" can include calculating, computing, processing, deriving, investigating, searching (e.g., searching within a table, database, or another data structure), verifying, etc. Further, "determining" can include receiving (e.g., receiving information), accessing (e.g., accessing data in a memory), etc. Also, "determining" can include resolving, selecting, choosing, establishing, etc.
[0197]
[0142] The methods disclosed herein comprise one or more steps or actions for achieving the methods. The steps and / or actions of the methods may be exchanged with each other without departing from the scope of the claims. In other words, unless a particular order of steps or actions is specified, the order and / or use of particular steps and / or actions may be changed without departing from the scope of the claims. Further, the various operations of the methods described above may be implemented by any suitable means capable of performing the corresponding functions. Those means may include various (one or more) hardware and / or software components and / or modules, including but not limited to circuits, application specific integrated circuits (ASICs), or processors. Generally, where there are operations shown in the figures, those operations may have corresponding means-plus-function components with similar numbers.
[0198]
[0143] The following claims are not intended to be limited to the embodiments shown herein but are to be accorded the full scope consistent with the language of the claims. References in the claims to an element in the singular are not intended to mean "one and only one" unless explicitly so stated, but rather "one or more." The term "some," unless otherwise specified, refers to one or more. A claim element is not to be construed under the provisions of 35 U.S.C. § 112(f) unless the element is expressly recited using the phrase "means for" or, in the case of a method claim, the element is recited using the phrase "step for." All structural and functional equivalents to the various aspects of the elements described throughout this disclosure that are known or later come to be known to those of ordinary skill in the art are expressly incorporated herein by reference and are intended to be encompassed by the claims. Further, nothing disclosed herein is intended to be dedicated to the public regardless of whether such disclosure is expressly recited in the claims. The invention described in the claims of the present application at the time of filing is appended below. [C1] A method for performing collaborative learning of a machine learning model, comprising: receiving, in a device, from a server that manages collaborative learning of a global machine learning model, a subset of model elements from a set of model elements for the global machine learning model, and a set of gate probabilities, wherein each gate probability of the set of gate probabilities is associated with one model element of the subset of model elements; generating, by the device, a set of model update values based on training a local machine learning model based on the set of model elements and the set of gate probabilities; and transmitting the set of model update values from the device to the server. A method comprising the above steps. [C2] The method according to C1, wherein the subset of model elements comprises a subset of weights associated with edges connecting nodes in the global machine learning model. [C3] The set of model update values comprises: a set of weight gradients associated with the local machine learning model; and a set of gate probability gradients associated with the local machine learning model. The method according to C2. [C4] The set of model update values comprises: a set of weight gradients associated with the local machine learning model; and binary gate variable values associated with each weight gradient of the set of weight gradients. The method according to C2. [C5] The method according to C1, wherein the subset of model elements comprises a subset of nodes in the global machine learning model. [C6] The method according to C1, wherein the subset of model elements comprises a subset of channels in a convolutional filter of the global machine learning model. [C7] The method according to C1, further comprising receiving, in the device, from the server a final set of model elements, wherein the final set of model elements corresponds to a pruned global machine learning model. [C8] A processing system, comprising: a memory comprising computer-executable instructions; and a processor configured to execute the computer-executable instructions, the processing system being caused to: receive, from a server that manages collaborative learning of a global machine learning model, a subset of model elements from a set of model elements for the global machine learning model, Receiving a set of gate probabilities, wherein each gate probability of the set of gate probabilities is associated with one of the model elements of the subset of model elements; Generating a set of model update values based on training a local machine learning model based on the set of model elements and the set of gate probabilities; Sending a set of model update values to the server; One or more processors for causing the above to be performed; A processing system comprising. [C9] The processing system according to C8, wherein the subset of model elements comprises a subset of weights associated with edges connecting nodes in the global machine learning model. [C10] The set of model update values is A set of weight gradients associated with the local machine learning model; and A set of gate probability gradients associated with the local machine learning model. The processing system according to C9, comprising. [C11] The set of model update values is A set of weight gradients associated with the local machine learning model; and A binary gate variable value associated with each weight gradient of the set of weight gradients. The processing system according to C9, comprising. [C12] The processing system according to C8, wherein the subset of model elements comprises a subset of nodes in the global machine learning model. [C13] The processing system according to C8, wherein the subset of model elements comprises a subset of channels in the convolutional filter of the global machine learning model. [C14] The one or more processors are further configured to receive a final set of model elements from the server, wherein the final set of model elements corresponds to a pruned global machine learning model. The processing system according to C8. [C15] A method for performing collaborative learning of a machine learning model, comprising: For each respective client of a plurality of clients and for each training round of a plurality of training rounds, The server generates a subset of model elements for each respective client based on sampling a gate probability distribution for each model element of a set of model elements for a global machine learning model; and From the server to each respective client, The subset of model elements, and sending a set of gate probabilities based on the sampling, where each gate probability of the set of gate probabilities is associated with one of the model elements of the subset of model elements; receiving, at the server, a respective set of model update values from each of the plurality of clients; updating, by the server, the global machine learning model based on the respective sets of model update values from each of the plurality of clients; A method comprising. [C16] The method according to C15, wherein the subset of model elements comprises a subset of weights associated with edges connecting nodes in the global machine learning model. [C17] Each respective set of model update values comprises a set of weight gradients associated with a local machine learning model trained by the respective client, and a set of gate probability gradients associated with the local machine learning model trained by the respective client; The method according to C16. [C18] Each respective set of model update values comprises a set of weight gradients associated with a local machine learning model trained by the respective client, and a binary gate variable value associated with each weight gradient of the set of weight gradients; The method according to C16. [C19] The method according to C15, wherein the subset of model elements comprises a subset of nodes in the global machine learning model. [C20] The method according to C15, wherein the subset of model elements comprises a subset of channels in the convolutional filter of the global machine learning model. [C21] The method according to C15, wherein updating the global machine learning model based on the respective sets of model update values from each of the plurality of clients by the server further comprises pruning the updated global machine learning model based on updated gate probabilities for the global machine learning model and a threshold gate probability value. [C22] A processing system comprising a memory comprising computer-executable instructions; and configured to execute the computer-executable instructions to cause the processing system to, for each of the plurality of clients and for each of a plurality of training rounds, Generating a subset of model elements for each of the respective clients based on sampling a gate probability distribution for each model element of a set of model elements for a global machine learning model; For each of the respective clients; Sending the subset of model elements and A set of gate probabilities based on the sampling, where each gate probability of the set of gate probabilities is associated with one model element of the subset of model elements; Receiving, from each respective client of the plurality of clients, a respective set of model update values; Updating the global machine learning model based on the respective sets of model update values from each respective client of the plurality of clients; One or more processors causing the above to be performed; A processing system comprising the same. [C23] The processing system according to C22, wherein the subset of model elements comprises a subset of weights associated with edges connecting nodes in the global machine learning model. [C24] Each respective set of model update values comprises: A set of weight gradients associated with a local machine learning model trained by each of the respective clients; and A set of gate probability gradients associated with the local machine learning model trained by each of the respective clients. The processing system according to C23. [C25] Each respective set of model update values comprises: A set of weight gradients associated with a local machine learning model trained by each of the respective clients; and A binary gate variable value associated with each weight gradient of the set of weight gradients. The processing system according to C23. [C26] The processing system according to C22, wherein the subset of model elements comprises a subset of nodes in the global machine learning model. [C27] The processing system according to C22, wherein the subset of model elements comprises a subset of channels in a convolutional filter of the global machine learning model. [C28] To update the global machine learning model based on each respective set of model update values from each of the plurality of clients, the one or more processors are further configured to prune the updated global machine learning model based on an updated gate probability for the global machine learning model and a threshold gate probability value, the processing system according to C22.
Claims
1. A method for performing collaborative learning of a machine learning model, comprising: at a device, receiving from a server that manages collaborative learning of a global machine learning model, a subset of model elements from a set of model elements for the global machine learning model, and a set of gate probabilities, wherein each gate probability of the set of gate probabilities is associated with one of the model elements of the subset of model elements, and the gate probability represents the likelihood that the associated model element is included in a local machine learning model for each device; generating, by the device, a set of model update values based on training the local machine learning model in the device based on the set of model elements and the set of gate probabilities, wherein the set of model update values comprises update values for model elements and update values for gate probabilities; transmitting, from the device to the server, the set of model update values for updating the global machine learning model and comprising a method.
2. The subset of model elements comprises a subset of weights associated with edges connecting nodes in the global machine learning model, The set of model update values is a set of weight gradients associated with the local machine learning model, and a set of gate probability gradients associated with the local machine learning model and comprising, or The set of model update values is a set of weight gradients associated with the local machine learning model, and a binary gate variable value associated with each weight gradient of the set of weight gradients and comprising the method according to claim 1.
3. The subset of model elements comprises a subset of nodes in the global machine learning model, or The subset of model elements comprises a subset of channels in a convolutional filter of the global machine learning model, the method according to claim 1.
4. The method according to claim 1, further comprising receiving, at the device, from the server a final set of model elements, wherein the final set of model elements corresponds to a pruned global machine learning model.
5. A processing system, comprising: a memory comprising computer-executable instructions, configured to execute the computer-executable instructions, causing the processing system to receive, from a server that manages federated learning of a global machine learning model, a subset of model elements from a set of model elements for the global machine learning model, and a set of gate probabilities, where each gate probability of the set of gate probabilities is associated with one of the model elements of the subset of model elements, and the gate probability represents the likelihood that the associated model element is included in a local machine learning model for each device generate a set of model update values based on training the local machine learning model at each device based on the set of model elements and the set of gate probabilities, the set of model update values comprising update values for model elements and update values for gate probabilities send the set of model update values for updating the global machine learning model to the server and one or more processors for causing the processing system to perform the above. **Claim 6** The subset of model elements comprises a subset of weights associated with edges connecting nodes within the global machine learning model, The set of model update values comprises a set of weight gradients associated with the local machine learning model, and a set of gate probability gradients associated with the local machine learning model, or The set of model update values comprises a set of weight gradients associated with the local machine learning model, and binary gate variable values associated with each weight gradient of the set of weight gradients The processing system according to claim 5. **Claim 7** The subset of model elements comprises a subset of nodes within the global machine learning model, or The subset of model elements comprises a subset of channels within a convolutional filter of the global machine learning model. The processing system according to claim 5. **Claim 8** The one or more processors are further configured to receive a final set of model elements from the server, where the final set of model elements corresponds to a pruned global machine learning model. The processing system according to claim 5. **Claim 9** A method for performing federated learning of a machine learning model, comprising For each of the plurality of clients and for each of the plurality of training rounds, Based on the server sampling a gate probability distribution for each model element of a set of model elements for the global machine learning model, generating a subset of model elements for each of the respective clients, where the gate probability represents the likelihood that the associated model element is included in the local machine learning model for each client, From the server to each of the respective clients, Sending the subset of model elements, And a set of gate probabilities based on the sampling, where each gate probability of the set of gate probabilities is associated with one model element of the subset of model elements, In the server, receiving, from each of the plurality of clients, a respective set of model update values, where the set of model update values comprises an update value for a model element and an update value for a gate probability, Updating the global machine learning model by the server based on the respective sets of model update values from each of the plurality of clients Comprising a method.
10. The subset of model elements comprises a subset of weights associated with edges connecting nodes in the global machine learning model, Each of the respective sets of model update values, A set of weight gradients associated with the local machine learning model trained by each of the respective clients, And a set of gate probability gradients associated with the local machine learning model trained by each of the respective clients Or, Each of the respective sets of model update values, A set of weight gradients associated with the local machine learning model trained by each of the respective clients, And a binary gate variable value associated with each weight gradient of the set of weight gradients The method according to claim 9.
11. The subset of model elements comprises a subset of nodes in the global machine learning model, or, The subset of model elements comprises a subset of channels in the convolutional filter of the global machine learning model, the method according to claim 10.
12. Updating the global machine learning model based on each respective set of model update values from each of the plurality of clients by the server further comprises pruning the updated global machine learning model based on an updated gate probability for the global machine learning model and a threshold gate probability value, the method of claim 10.
13. A processing system, A memory comprising computer-executable instructions, Configured to execute the computer-executable instructions, the processing system, For each respective client of the plurality of clients and for each training round of the plurality of training rounds, Generating a subset of model elements for each respective client based on sampling a gate probability distribution for each model element of a set of model elements for a global machine learning model, wherein a gate probability represents the likelihood that an associated model element is included in a local machine learning model for a respective client, To each respective client, Transmitting a subset of the model elements and A set of gate probabilities based on the sampling, wherein each gate probability of the set of gate probabilities is associated with one of the model elements of the subset of model elements, Receiving from each respective client of the plurality of clients a respective set of model update values, the set of model update values comprising an update value for a model element and an update value for a gate probability, Updating the global machine learning model based on each respective set of model update values from each of the plurality of clients One or more processors for causing A processing system comprising.
14. The subset of model elements comprises a subset of weights associated with edges connecting nodes in the global machine learning model, Each respective set of model update values, A set of weight gradients associated with a local machine learning model trained by each respective client and A set of gate probability gradients associated with the local machine learning model trained by each respective client comprising, or each of the sets of model update values is a set of weight gradients associated with a local machine learning model trained by each of the respective clients, and a binary gate variable value associated with each weight gradient of the set of weight gradients The processing system according to claim 13, comprising.
15. The one or more processors are further configured to prune the updated global machine learning model based on an updated gate probability for the global machine learning model and a threshold gate probability value to update the global machine learning model based on each of the respective sets of model update values from each of the respective clients of the plurality of clients. The processing system according to claim 13.
Citation Information
Patent Citations
Customized identifiers across common features
JP2017520825A
Communication Efficient Federated Learning
US20200242514A1