Model training system, device, method, and program

The model learning system improves communication efficiency and learning accuracy in federated learning by averaging update differences and using a control variable to correct data bias, overcoming the challenges faced by existing methods like FedAvg and SCAFFOLD.

WO2025126302A1PCT designated stage expired Publication Date: 2025-06-19NT T INC
View PDF 0 Cites 0 Cited by

Patent Information

Application Number
PCT/JP2023/044348
Authority / Receiving Office
WO · WO
Patent Type
Applications
Current Assignee / Owner
Filing Date
2023-12-12
Publication Date
2025-06-19

AI Technical Summary

Technical Problem

Existing federated learning methods, such as model averaging FedAvg and SCAFFOLD, face challenges in achieving high learning accuracy and efficiency, particularly due to data bias and increased communication volume.

Method used

A model learning system that includes a server device and multiple model learning devices, which employs a selection unit to choose a subset of devices, a reception unit to gather update differences, a model parameter update unit to average these updates, and a transmission unit to disseminate the updated model parameters, while also using a control variable to correct data bias and optimize communication efficiency.

Benefits of technology

The proposed system enhances learning efficiency in communication by reducing the communication volume while maintaining convergence and accuracy, thereby addressing the limitations of previous methods.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure JP2023044348_19062025_PF_FP_ABST
    Figure JP2023044348_19062025_PF_FP_ABST
Patent Text Reader

Abstract

A model training system according to one aspect of the disclosed technology includes a server device and N model training devices. The server device 10 comprises: (1) a selection unit that selects a set K' from the N model training devices; (2) a reception unit that receives an update difference ut k from each model training device k belonging to the set K'; (3) a model parameter update unit that obtains wt using wt-1 and the update difference ut k; and (4) a transmission unit that transmits wt to each of the N model training devices. Each model training device k comprises: (1) a reception unit that receives wt-1 from the server device 10; (2) a control variable update unit that, if t>1, obtains λt-1 using at least wt-1 and λt-2; (3) a correction gradient calculation unit that obtains a correction gradient using wt-1, λt-1, and training data retained by each model training device k; (4) a clipping unit that adjusts the correction gradient to be equal to or less than a prescribed magnitude to obtain an adjusted correction gradient; and (5) a transmission unit that transmits an update difference obtained from the adjusted correction gradient to the server device 10.
Need to check novelty before this filing date? Find Prior Art

Description

Model learning system, device, method and program

[0001] The disclosed technology relates to federated learning.

[0002] Federated learning allows for the creation of a model that learns the characteristics of all data while protecting the data of participants in federated learning by sharing only the model (learning gradient) without sharing data. One well-known federated learning method is the model average (FedAvg).

[0003] However, when each participant is learning from different data, it is known that the model average FedAvg method does not allow learning to progress and does not improve learning accuracy.

[0004] Therefore, a method called SCAFFOLD has been proposed as a method for performing associative learning with high accuracy even in the presence of data bias by controlling the variance of the main variable (i.e., the model to be learned) using a control value (see, for example, Non-Patent Document 1).

[0005] Karimireddy, SP, Kale, S., Mohri, M., Reddi, S., Stich, S., Suresh, AT. (2020). SCAFFOLD: Stochastic Controlled Averaging for Federated Learning. Proceedings of the 37th International Conference on Machine Learning, in Proceedings of Machine Learning Research 119:5132-5143.

[0006] However, in SCAFFOLD, although the learning convergence is good, the communication volume is about twice that of the model average FedAvg, which means that the learning efficiency relative to communication is poor.

[0007] The disclosed technology aims to provide a model learning system, device, method, and program that have better learning efficiency for communications than conventional systems.

[0008] A model learning system according to one aspect of the disclosed technology is a model learning system including a server device and N model learning devices, where N is a predetermined positive integer, K is a predetermined positive integer equal to or less than N, k is an index indicating one of the model learning devices included in a set K′ consisting of K model learning devices among the N model learning devices, k=1,...,K, t is the number of rounds, and T d is a given positive integer, t=1,…,T d and w t are the parameters of the model that model learning device k is learning in round t, w0 is the initial value of the parameters of the model that model learning device k is trying to learn, and λ t is a control variable for correcting the bias in the data of model learning device k, and λ is a predetermined initial value. The server device includes: (1) a selection unit that selects set K′ from N model learning devices; and (2) an update difference u from each model learning device k belonging to set K′. t k a receiving unit that receives (3) w t-1 and update diff u t k Using w t and (4) a model parameter updater for calculating w t a transmitter for transmitting (1) w to each of the N model learning devices, and each model learning device k transmits (1) w t-1 (2) if t>1, then w t-1 and λ t-2 and λ t-1 (3) a control variable update part to calculate w t-1 and λ t-1and (3) a correction gradient calculation unit that calculates a correction gradient using the learning data held by each model learning device k, (4) a clipping unit that calculates an adjusted correction gradient by adjusting the correction gradient so that it is equal to or smaller than a predetermined magnitude, and (5) a transmission unit that transmits to the server device an update difference calculated from the adjusted correction gradient, wherein the processing of the selection unit of the server device, the reception unit of the server device, the model parameter update unit of the server device, the transmission unit of the server device, the reception unit of each model learning device k, the control variable update unit of each model learning device k, the correction gradient calculation unit of each model learning device k, the clipping unit of each model learning device k, and the transmission unit of each model learning device k is performed for t=1, ..., T d This is done for each of the above.

[0009] According to the disclosed technology, model learning with better learning efficiency for communication than conventional techniques is possible.

[0010] FIG. 1 is a diagram showing an example of the functional configuration of a model learning system. FIG. 2 is a diagram showing an example of the functional configuration of a server device. FIG. 3 is a diagram showing an example of a processing procedure of a server device. FIG. 4 is a diagram showing an example of the functional configuration of a model learning device. FIG. 5 is a diagram showing an example of a processing procedure of a model learning device. FIG. 6 is a diagram showing an example of processing to be executed by a model learning system and method. FIG. 7 is a diagram showing an example of the functional configuration of a computer.

[0011] Hereinafter, embodiments of the disclosed technology will be described with reference to the drawings. Note that components having the same functions in the drawings are given the same reference numerals, and redundant description will be omitted.

[0012] The symbol "^" used in the text should be written directly above the character immediately following it, but due to limitations in text notation, it is written immediately before the character in question.

[0013] As shown in FIG. 1, the model learning system includes a server device 10 and N model learning devices 1, ..., N. N is a predetermined positive integer. The N model learning devices 1, ..., N are connected to the server device 10 so as to be able to send and receive information. Let k be an index indicating one of the model learning devices included in a set K' consisting of K model learning devices among the N model learning devices. k = 1, ..., K.

[0014] The model learning method is realized, for example, by each component of the model learning system performing the processes of steps S101 to S105 and steps Sk1 to Sk5 shown in Figures 3 and 5. Specifically, the model learning method is realized, for example, by each component of the server device 10 performing the processes of steps S101 to S105 shown in Figure 3, and each component of the model learning device k performing the processes of steps Sk1 to Sk5 shown in Figure 5. Note that the processes of steps Sk1 to Sk5 are performed before the process of step S102 is performed.

[0015] As shown in FIG. 2, the server device 10 includes a selection unit 101, a reception unit 102, a model parameter update unit 103, a transmission unit 104, and a control unit 105.

[0016] t is the number of rounds, T d is a given positive integer, t=1,…,T d The server device 10 performs the following processing of steps S101 to S104 for each round t. The repeated processing of steps S101 to S104 is performed under the control of the control unit 105, which will be described later. As mentioned above, before the processing of step S102 is performed, the processing of steps Sk1 to Sk5 is performed. For this reason, under the control of the control unit 105, the processing of the selection unit of the server device 10, the reception unit of the server device 10, the model parameter update unit of the server device 10, the transmission unit of the server device 10, the reception unit k1 of each model learning device k, the control variable update unit k2 of each model learning device k, the correction gradient calculation unit k3 of each model learning device k, the clipping unit k4 of each model learning device k, and the transmission unit k5 of each model learning device k is performed for t=1, ..., T d This will be carried out for each of the above.

[0017] The selection unit 101 selects a set K′ from N model learning devices (step S101). The selection unit 101 selects the set K′ at random, for example.

[0018] The process in step S101 corresponds to "4:" in FIG.

[0019] The receiving unit 102 receives the update difference u from each model learning device k belonging to the set K′. t k The received update difference u is received (step S102). t k is transmitted to the model parameter update unit 103.

[0020] update difference u t k is obtained by the processing of steps Sk1 to Sk5 of the model learning device k, which will be described later.

[0021] The process in step S102 corresponds to "6:" in FIG.

[0022] The model parameter update unit 103 updates w t-1 and update diff u t k Using w t is calculated (step S103).

[0023] Specifically, the model parameter update unit 103 updates w t =w t-1 +(1 / |K'|)Σ k u t k That is, the model parameter update unit 103 calculates w t-1 The update difference u from each model learning device k belonging to the set K' is t k Add the average value of t |K'| is the number of model learning devices that belong to set K'.

[0024] The model parameter update unit 103 updates w t =w t-1 +Σ k α k u t k For k=1,...,K, α k is Σ k α k = 1. That is, the model parameter update unit 103 updates w t-1 The update difference u from each model learning device k belonging to the set K' is t kThe weighted average of t It may also be possible to use the following.

[0025] The process in step S103 corresponds to "8:" in FIG.

[0026] w t is the parameter of the model that model learning device k is learning in round t. w0 is the initial value of the parameter of the model that model learning device k is trying to learn. w t , w0 is a vector with the same number of dimensions as the number of dimensions D of the neural network, which is the model to be trained. w0 is determined randomly, for example. In the example of Figure 6, the process of setting w0 is performed in "2:".

[0027] The transmitting unit 104 receives the signal w t is transmitted to each of the N model learning devices (step S104).

[0028] The process in step S104 corresponds to "10:" in FIG.

[0029] Next, the configuration of model learning device k will be described, where k = 1, ..., K. In other words, the configuration of the model learning devices included in set K' will be described. Note that model learning devices other than the model learning devices included in set K' among model learning devices 1, ..., N also have the same configuration.

[0030] As shown in FIG. 4, the model learning device k includes, for example, a receiving unit k1, a control variable updating unit k2, a correction gradient calculating unit k3, a clipping unit k4, and a transmitting unit k5.

[0031] The receiver k1 is t-1 is received from the server device 10 (step Sk1). t is output to the control variable update unit k2 and the correction gradient calculation unit k3. t may be output to the transmitter k5. t-1 is the w transmitted by the transmitting unit 104 one round ago. t is.

[0032] When t>1, the control variable update unit k2 updates w t-1 and λ t-2and λ t-1 is calculated (step Sk2).

[0033] t=1,…,T d As, λ t is a control variable for correcting bias in the data of model learning device k. λ0 is a predetermined initial value.

[0034] In this way, the model learning device k can calculate the control variable λ based only on the information it has. t By calculating the control variable λ t In this respect, the learning efficiency with respect to communication is improved compared to the conventional method.

[0035] For example, when t>1, the control variable update unit k2 updates λ t-2 +w t-1 -w' t-2 Calculate the result as λ t-1 Let w' t-2 is calculated in "24:" of Figure 6 two rounds ago. t is.

[0036] The control variable update unit k2 initializes λ0 when t = 1. For example, the control variable update unit k2 randomly sets λ0.

[0037] The process of step Sk2 corresponds to "14:" to "18:" in FIG.

[0038] The correction gradient calculation unit k3 is t-1 and λ t-1 and the learning data held by each model learning device k, λ t-1 A corrected gradient is calculated (step Sk3).

[0039] For example, the correction gradient calculation unit k3 calculates the gradient using a predetermined optimization method. Examples of the optimization method include stochastic optimization methods such as AdamW, AdaGrad, RMSprop, SGD, MomentumSGD, and Momentum Averaging. Of course, other optimization methods may also be used.

[0040] The control variables are used to correct the gradients calculated by each model learning device. Furthermore, the control variables themselves are updated using the differences in the model parameters between the model learning devices. Therefore, any optimization method can be used, regardless of the control variables.

[0041] Specifically, Δw t G =AdamW(G,w',T gd )-w t-1 -λ t-1 Calculate the result and use it as the correction gradient Δw t G Let's say.

[0042] Here, G is any set included in the subset M. The set D of training data held by the model training device k is k Let G be a set of disjoint sets whose elements are the training data contained in k As such, M is a set G k M is a subset of M. M is selected at random, for example. In the example of FIG. 6, the selection of M is performed by "13:".

[0043] For example, the data held by the model learning device k is expressed as d1, ..., d 100 Then, D k ={d1,…,d 100} In this case, for example, G k ={{d1,…,d 10}, {d 11 ,…,d 20},…,{d 91 ,…,d 100}}. In this case, for example, M is {{d1,...,d 10}, {d 21 ,…,d 30}}. In this example, the set G k The index k of the training data d included in the set that is an element of G is a consecutive integer. k The index k of the training data d included in the set that is an element of G may be random. For example, k One of the element sets of {d2, d 19 ,…,d 70} may also be used.

[0044] w'=w t-1 T gd is the number of times to perform gradient descent, which is a predetermined positive integer.

[0045] The process of step Sk3 corresponds to "21:" in FIG.

[0046] The clipping unit k4 adjusts the correction gradient so that it is equal to or smaller than a predetermined magnitude, and obtains an adjusted correction gradient (step Sk4).

[0047] For example, the clipping unit k4 has a correction gradient Δw t G Using Δw t G / max(1,||Δw t G ||2 / S) and use the result as the adjusted correction gradient Δ^w t G Let ||Δw t G ||2 is the correction gradient Δw t G is the L2 norm of

[0048] The process of step Sk4 corresponds to "22:" in FIG.

[0049] The transmitter k5 receives the adjusted correction gradient Δ̂w t G The update difference calculated from the above is transmitted to the server device 10 (step Sk5).

[0050] In the example of FIG. 6, from “19:” to “23:”, the adjusted corrected gradient Δ̂w is calculated from part of the learning data held by the model learning device k. t G In other words, the correction gradient calculation unit k3 and the clipping unit k4 may perform the process of determining the adjusted correction gradient multiple times based on part of the learning data held by the model learning device k.

[0051] In this case, the update difference calculated from the adjusted correction gradient may be the average value of multiple correction gradients calculated based on part of the learning data held by each model learning device k. In other words, the update difference calculated from the adjusted correction gradient is (1 / |M|)(Σ G Δ^w t G ) where |M| is the number of elements in set M.

[0052] Furthermore, the update difference calculated from the adjusted correction gradient may be a weighted average value of a plurality of correction gradients calculated based on a portion of the learning data held by each model learning device k. In other words, the update difference calculated from the adjusted correction gradient is Σ G β G Δ^w t G β G is Σ G β G =1, and is a predetermined value between 0 and 1.

[0053] The update difference is the adjusted correction gradient Δ^w t G It may be itself.

[0054] The process of step Sk5 corresponds to "25:" in Fig. 6. In "25:" in Fig. 6, the update difference is set to w' t -w t-1 As shown in "24:" in Figure 6, w' t =w t-1 -w t-1 +(1 / |M|)(Σ G Δ^w t G ) Therefore, w' t -w t-1 =(1 / |M|)(Σ G Δ^w t G ) In addition, the w' of "24:" in Figure 6 t The process of obtaining is performed by the model learning device k.

[0055] [Modifications] The specific configurations of the embodiments of the disclosed technology are not limited to the configurations described above. The specific configurations of the embodiments of the disclosed technology can be appropriately modified in design, etc., within the scope of the spirit of the embodiments of the disclosed technology.

[0056] The various processes described in the embodiments of the disclosed technology may not only be performed chronologically in the order described, but may also be performed in parallel or individually depending on the processing capacity of the device performing the processes or as needed.

[0057] For example, data may be exchanged directly between components of the model learning device or via a storage unit (not shown).Furthermore, data may be exchanged directly between model learning devices or via a relay device (not shown), such as an aggregation server.

[0058] It goes without saying that other modifications are possible without departing from the spirit of the present invention.

[0059] All publications, patent applications, and technical standards mentioned in this specification are herein incorporated by reference to the same extent as if each individual publication, patent application, or technical standard was specifically and individually indicated to be incorporated by reference.

[0060] [Program, Recording Medium] The functions realized by the components described in this specification may be implemented in circuitry or processing circuitry, including general-purpose processors, application-specific processors, integrated circuits, ASICs (Application Specific Integrated Circuits), CPUs (Central Processing Units), conventional circuits, and / or combinations thereof, programmed to realize the described functions. A processor includes transistors and other circuits and is considered to be circuitry or processing circuitry. A processor may also be a programmed processor that executes a program stored in a memory.

[0061] In this specification, a circuitry, unit, or means is hardware that is programmed to realize or performs the described functions, which may be any hardware disclosed herein or any hardware known to be programmed to realize or perform the described functions.

[0062] If the hardware is a processor considered to be a type of circuitry, the circuitry, means, or unit is a combination of the hardware and software used to configure the hardware and / or processor.

[0063] The various processes described above can be implemented by loading a program that executes each step of the above method into the recording unit 2020 of the computer 2000 shown in Figure 7, and operating the control unit 2010, input unit 2030, output unit 2040, display unit 2050, etc.

[0064] The program describing the processing contents can be recorded on a computer-readable recording medium, which may be, for example, a magnetic recording device, an optical disk, a magneto-optical recording medium, a semiconductor memory, or any other suitable recording medium.

[0065] The program may be distributed by, for example, selling, transferring, lending, etc. portable recording media such as DVDs and CD-ROMs on which the program is recorded. Furthermore, the program may be stored in a storage device of a server computer, and then transferred from the server computer to other computers via a network, thereby distributing the program.

[0066] A computer that executes such a program may first temporarily store the program recorded on a portable recording medium or transferred from a server computer in its own storage device. Then, when executing a process, the computer reads the program stored on its own recording medium and executes the process in accordance with the read program. Alternatively, the computer may read the program directly from a portable recording medium and execute the process in accordance with the program. Furthermore, the computer may execute the process in accordance with the program each time a program is transferred from a server computer to the computer. Alternatively, the server computer may not transfer the program to the computer, but may instead execute the process through a so-called ASP (Application Service Provider) service, which realizes the processing function by issuing an execution instruction and obtaining the results. Furthermore, the server computer may execute the process at the terminal using a so-called SaaS (Software as a Service) service, which allows users to use part of a server computer along with the program. In this embodiment, the program includes information used for processing by an electronic computer that is equivalent to a program (such as data that is not a direct instruction to a computer but has properties that dictate computer processing).

[0067] Furthermore, in this embodiment, the device is configured by executing a predetermined program on a computer, but at least a part of the processing contents may be realized by hardware.

Claims

1. A model learning system including a server device and N model learning devices, where N is a predetermined positive integer, K is a predetermined positive integer less than or equal to N, k is an index indicating any one of the K model learning devices included in a set K' composed of K of the N model learning devices, k = 1,..., K, t is the number of rounds, and T d is a predetermined positive integer, t = 1,..., T d and w t is the parameter of the model that the model learning device k is learning in round t, w0 is the initial value of the parameter of the model that the model learning device k is going to learn, and λ t is a control variable for correcting the bias of the data of the model learning device k. Assuming that λ0 is a predetermined initial value, the server device includes: (1) a selection unit that selects the set K' from the N model learning devices; (2) a reception unit that receives the update difference u t k from each model learning device k belonging to the set K'; (3) a model parameter update unit that obtains w t-1 using w t k and the update difference u t ; and (4) a transmission unit that transmits w t to each of the N model learning devices. Each model learning device k includes: (1) a reception unit that receives w t-1 from the server device; (2) when t > 1, a control variable update unit that obtains λ t-1 using at least w t-2 and λ t-1 ; and (3) w t-1 and λ t-1a correction gradient calculation unit that obtains a correction gradient using the learning data possessed by each of the model learning devices k; (4) a clipping unit that obtains an adjusted correction gradient obtained by adjusting the correction gradient to a predetermined magnitude or less; and (5) a transmission unit that transmits an update difference obtained from the adjusted correction gradient to the server device, and the selection unit of the server device, the reception unit of the server device, the model parameter update unit of the server device, the transmission unit of the server device, the reception unit of each model learning device k, the control variable update unit of each model learning device k, the correction gradient calculation unit of each model learning device k, the clipping unit of each model learning device k, and the processing of the transmission unit of each model learning device k are performed for each of t = 1, …, T d respectively, a model learning system.

2. The model learning system according to claim 1, wherein the correction gradient calculation unit and the clipping unit perform a process of obtaining the adjusted correction gradient a plurality of times based on a part of the learning data possessed by each model learning device k, and the update difference is an average value of a plurality of correction gradients obtained based on a part of the learning data possessed by each model learning device k. A model learning system.

3. t is the number of rounds, and T d is a predetermined positive integer, and t = 1, …, T d where w t is the parameter of the model being learned by the model learning device in round t, w0 is the initial value of the parameter of the model that the model learning device k intends to learn, and λ t is a control variable for correcting the bias of the data of the model learning device, and assuming that λ0 is a predetermined initial value, a reception unit that receives w t-1 from the server device; and when t > 1, at least w t-1 and λ t-2 are used to obtain λ t-1 a control variable update unit; w t-1 and λ t-1A correction gradient calculation unit that obtains a correction gradient using the learning data of the model learning device, a clipping unit that obtains an adjusted correction gradient obtained by adjusting the correction gradient to be equal to or less than a predetermined magnitude, and a transmission unit that transmits an update difference obtained from the adjusted correction gradient to the server device, wherein the processing of the reception unit, the control variable update unit, the correction gradient calculation unit, the clipping unit, and the transmission unit is performed for t = 1, …, T d respectively, and a model learning device.

4. A model learning method performed by a model learning system including a server device and N model learning devices, where N is a predetermined positive integer, K is a predetermined positive integer less than or equal to N, k is an index indicating any one of the K model learning devices included in a set K' composed of K of the N model learning devices, k = 1, …, K, t is the number of rounds, and T d is a predetermined positive integer, and t = 1, …, T d where w t is the parameter of the model being learned by model learning device k in round t, w0 is the initial value of the parameter of the model that model learning device k is about to learn, λ t is a control variable for correcting the bias of the data of model learning device k, and λ0 is a predetermined initial value. The selection unit of the server device selects the set K' from the N model learning devices. The reception unit of the server device receives an update difference u t k from each model learning device k belonging to the set K'. The model parameter update unit of the server device obtains w t-1 using w t k and the update difference u t . The transmission unit of the server device transmits w t to each of the N model learning devices. The reception unit of each model learning device k receives w t-1 from the server device. When t > 1, the control variable update unit of each model learning device k uses w t-1 and λ t-2using at least to obtain λ t-1 a step of obtaining, and for each correction gradient calculation unit of the model learning device k, w t-1 and λ t-1 and the learning data of each model learning device k to obtain a correction gradient; a step for each clipping unit of the model learning device k to obtain an adjusted correction gradient obtained by adjusting the correction gradient to be equal to or less than a predetermined magnitude; a step for each transmission unit of the model learning device k to transmit an update difference obtained from the adjusted correction gradient to the server device; and the processing of each step is performed for each of t = 1,..., T d A model learning method including 5. t is the number of rounds, and T d is a predetermined positive integer, t = 1,..., T d and w t is the parameter of the model being learned by the model learning device in round t, w0 is the initial value of the parameter of the model that the model learning device k intends to learn, and λ t is a control variable for correcting the bias of the data of the model learning device, and assuming that λ0 is a predetermined initial value, a step for the receiving unit to receive w t-1 from the server device; a step for the control variable update unit to obtain λ t-1 using at least w t-2 and λ t-1 when t > 1; a step for the correction gradient calculation unit to obtain a correction gradient using w t-1 and λ t-1 and the learning data it has; a step for the clipping unit to obtain an adjusted correction gradient obtained by adjusting the correction gradient to be equal to or less than a predetermined magnitude; a step for the transmission unit to transmit an update difference obtained from the adjusted correction gradient to the server device, and the processing of each step is performed for each of t = 1,..., T d A model learning method 6. A program for causing a computer to execute each step of the communication method according to claim 5.