A method for training a segmentation federated learning model based on a heterogeneous system
By personalized segmentation and parallel training of client models in the segmentation learning framework, the problem of inefficient training of heterogeneous devices is solved, and more efficient model training and collaborative computing are achieved.
Patent Information
- Application Number
- CN202411728336.X
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2024-11-28
- Publication Date
- 2025-07-01
- Estimated Expiration
- 2044-11-28
AI Technical Summary
The current segmented learning framework faces the problem of inefficient training due to its serial training method. When heterogeneous edge devices run the same client model, due to the differences in computing and communication resources of different devices, the slowest devices will slow down the training efficiency of the entire system.
By personalizing the client model to adapt to the conditions of the client device, the calculation burden of heterogeneous devices is reduced, and the batch size is dynamically adjusted to adapt to the computing power and channel conditions of the device.
This improves the model training efficiency of heterogeneous systems, reduces the total training delay, ensures the coordinated training of heterogeneous devices, and improves the overall computing efficiency of distributed systems.
Smart Images

Figure CN119312947B_ABST
Abstract
Description
Technical Field
[0001] The present invention relates to the technical fields of distributed systems and split federated learning, and particularly to a method for training a split federated learning model based on a heterogeneous system. Background Art
[0002] In the fields of artificial intelligence and the Internet of Things, data interconnection and collaborative computing among intelligent devices have been widely applied. Against this background, distributed learning architectures have emerged, especially federated learning (FL) and split learning (SL), which have become key ways to help resource-constrained devices achieve efficient model training. Split learning cleverly divides the model into two parts, a client side and a server side, avoiding the transmission of raw data, thus protecting user privacy and significantly reducing the computational burden on devices, and further improving the operation efficiency of the entire distributed system.
[0003] The current split learning framework faces the problem of low training efficiency due to its serial training method. At the same time, when heterogeneous edge devices run the same client model, due to differences in the computing and communication resources of different devices, the slowest device will slow down the training efficiency of the entire system. Summary of the Invention
[0004] In view of this, embodiments of the present invention provide a method for training a split federated learning model based on a heterogeneous system, which performs personalized splitting of the client model according to the conditions of the client device, reduces the computational burden on heterogeneous devices, and improves the model training efficiency of the heterogeneous system.
[0005] One aspect of the present invention provides a method for training a split federated learning model based on a heterogeneous system. The heterogeneous system includes a plurality of clients, a central server, and an edge server; the global model is split into client local models adapted to each client and corresponding server local models based on the device conditions of each client; the client local models are used to be deployed on the corresponding clients, and the server local models are used to be deployed on the central server; wherein, the device conditions of each client include the device storage resources, computing frequency, and channel conditions of the client, and based on the historical data of the operation of the heterogeneous system, the optimal splitting point for each client to split the global model is determined with the goal of minimizing the total training delay; the training method includes the following steps:
[0006] In each round of training, all clients execute the training of the client local models in parallel. The client local models deployed on each client perform forward calculation using the training sample set constructed from the local data of each client and output shredded data;
[0007] Each client transmits the shredded data and the labels of the training sample set to the server local model of the central server for forward calculation, and outputs a predicted value; calculates the loss value corresponding to each client based on the predicted value and the label;
[0008] Calculates the gradient of the output layer of each server local model according to the loss value, performs backpropagation of the gradient in the server local model, obtains the gradient of the shredded data, and transmits it back to the corresponding client;
[0009] The client local model of each client performs backpropagation of the gradient of the shredded data to update the parameters of the client local model;
[0010] After all client local models are updated, transmit all the updated client local models to the edge server;
[0011] Set the network layer in the global model that contains the corresponding split points of all clients in the global model as the common layer; the edge server aggregates the parameters of the common layer involved in all updated client local models, and the central server aggregates the parameters of the common layer involved in all server local models after all backpropagations are completed, and exchanges information on the parameters of the common layer aggregated on both sides; the edge server aggregates all client local models after information exchange to obtain a global client model; the central server aggregates all server local models after information exchange to obtain a global server model;
[0012] The edge server distributes the client local models to the corresponding clients according to the split points at which each client splits the global model for the global client model.
[0013] In some embodiments of the present invention, the method further includes:
[0014] In each round of training, divide the training sample set of each client into multiple mini-batch data, and use the mini-batch gradient descent algorithm to iteratively update the parameters of this round of training.
[0015] In some embodiments of the present invention, for the batch size setting of the mini-batch data:
[0016] Set the minimum batch size and the maximum batch size of the mini-batch data; each time the training sample set is divided during training, dynamically adjust the batch size between the minimum batch size and the maximum batch size according to the current computing power and channel conditions of each client; the batch size of the mini-batch data is positively correlated with the computing power and the quality of the channel conditions.
[0017] In some embodiments of the present invention, determining the optimal splitting point for each client to split the global model with the goal of minimizing the total training delay is based on optimizing the single-round training delay of each client. The single-round training delay of each client is as follows:
[0018]
[0019] Wherein, represents the forward propagation delay of the client's local model, represents the forward propagation computation volume, and f k represents the computing frequency of the client, and κ k represents the computing density of the client;
[0020] represents the fragmented data transmission delay, and bξ s (l k ) represents the size of the fragmented data, represents the uplink transmission rate from the client to the central server;
[0021] represents the forward propagation delay of the model on the central server, represents the computing frequency allocated by the central server to client k, and κ s represents the computing density of the central server;
[0022] represents the backpropagation delay of the model on the central server;
[0023] represents the gradient backpropagation delay, and bξ g (l k ) represents the size of the gradient of the fragmented data;
[0024] represents the backpropagation delay of the client's local model;
[0025] represents the delay for the client's local model to be transmitted to the edge server, and ξ m (l k ) represents the size of the client's local model, represents the uplink transmission rate from the client to the edge server;
[0026] represents the delay for the edge server to distribute the client's local model;
[0027] N represents the number of mini-batch data into which the training sample set is divided, obtained by dividing the training sample set by the batch size of the mini-batch data.
[0028] In some embodiments of the present invention, before each round of training, the method further includes:
[0029] Based on each client's local model determined by the optimal segmentation point, with the goal of minimizing the total training delay, determine the optimal computing frequency allocated by the central server to each client, the optimal uplink and downlink transmission powers from each client to the central server, and the optimal uplink and downlink transmission powers from each client to the edge server.
[0030] In some embodiments of the present invention, the method for the edge server to aggregate the global client model and the central server to aggregate the global server model is: weighted average the parameters of each layer of the model with the proportion of the local data volume of each client in the total data volume of all clients as the weight to obtain the aggregated model parameters.
[0031] In some embodiments of the present invention, the method further includes:
[0032] In a single training round, the edge server regularly sends signals to each client to confirm the online status of each client. If no reply message is received within the preset time, mark this client as an abnormal status;
[0033] For the clients marked as abnormal status in the current training round, the edge server abandons waiting for these clients to transmit the updated client local models and only aggregates the local models of all clients except those with abnormal status.
[0034] In some embodiments of the present invention, encrypted data transmission is performed between the client, the edge server, and the central server.
[0035] Another aspect of the present invention provides a computer-readable storage medium, on which computer programs / instructions are stored. When the computer programs / instructions are executed by a processor, the steps of the method described in any one of the above are implemented.
[0036] Another aspect of the present invention further provides a computer program product, including computer programs / instructions. When the computer programs / instructions are executed by a processor, the steps of the method described in any one of the above are implemented.
[0037] The beneficial effects of the present invention are at least:
[0038] The present invention discloses a method for training a split federated learning model based on a heterogeneous system, which includes multiple clients, a central server, and an edge server. The global model divides the client local model and the server local model based on each client device condition and deploys them to the corresponding client and the central server respectively. In each round of training, the clients train in parallel, forward-propagate the local data to output shredded data and transmit it to the server local model of the central server to continue forward propagation to obtain a loss value, and then back-propagate the gradient to update the parameters of the server local model and the client local model. Each client transmits the updated model to the edge server, and after the edge server and the central server complete the parameter exchange of the common layer, they conduct a model aggregation to generate a global client model and a global server model. The global client model distributes the client local model to the clients according to the segmentation points of each client. The present invention adapts to the conditions of client devices to perform personalized segmentation on the client models, reduces the computational burden of heterogeneous devices, and improves the model training efficiency of the heterogeneous system.
[0039] Additional advantages, objects, and features of the present invention will be partly described below, and will partly become apparent to those of ordinary skill in the art after studying the following part, or can be learned from the practice of the present invention. The objects and other advantages of the present invention can be achieved and obtained by the structure specifically pointed out in the specification and the drawings.
[0040] Those skilled in the art will understand that the objects and advantages that can be achieved by the present invention are not limited to the above specifically described, and the above and other objects that the present invention can achieve will be more clearly understood according to the following detailed description. Brief Description of the Drawings
[0041] The drawings described herein are used to provide a further understanding of the present invention, form a part of this application, and do not limit the present invention. In the drawings:
[0042] Figure 1 It is a flowchart of a method for training a split federated learning model based on a heterogeneous system in an embodiment of the present invention.
[0043] Figure 2 It is a framework diagram of heterogeneous split federated training in another embodiment of the present invention.
[0044] Figure 3 It is a structure diagram of the delay of heterogeneous split federated training in another embodiment of the present invention.
[0045] Figure 4 It is a call relationship diagram of a resource optimization algorithm in another embodiment of the present invention. Detailed Description of the Embodiments
[0046] To make the objectives, technical solutions, and advantages of the present invention more clearly understood, the present invention will be further described in detail below in conjunction with the embodiments and the accompanying drawings. Herein, the illustrative embodiments of the present invention and their descriptions are used to explain the present invention, but not to limit the present invention.
[0047] Herein, it should also be noted that in order to avoid obscuring the present invention due to unnecessary details, only the structures and / or processing steps closely related to the solution according to the present invention are shown in the drawings, while other details less relevant to the present invention are omitted.
[0048] It should be emphasized that the term "comprising / including" when used herein refers to the presence of features, elements, steps, or components, but does not exclude the presence or addition of one or more other features, elements, steps, or components.
[0049] Herein, it should also be noted that if not otherwise specified, the term "connection" in this document can refer not only to a direct connection, but also to an indirect connection with an intermediate.
[0050] In the following, embodiments of the present invention will be described with reference to the accompanying drawings. In the drawings, the same reference numerals represent the same or similar components, or the same or similar steps.
[0051] Federated learning is a distributed machine learning method that allows multiple participants (such as devices or organizations) to jointly train a global model without sharing their local data. Each participant uses its own local data for model training and sends updates (such as weight adjustments) to a central server without directly transmitting the data.
[0052] Split learning is a distributed learning approach that divides the model into two parts: a front end and a back end. The front-end model is executed on a local client, while the back-end model runs on a central server. The client calculates the output of the front-end model by inputting data and then sends this output (instead of the original data) to the server, which continues the processing to generate the final prediction.
[0053] An embodiment of the present invention provides a method for training a split federated learning model based on a heterogeneous system. The heterogeneous system includes multiple clients, a central server, and an edge server. The global model is split into client local models adapted to each client and corresponding server local models based on the device conditions of each client. The client local models are used to be deployed on the corresponding clients, and the server local models are used to be deployed on the central server. Among them, the device conditions of each client include the device storage resources, computing frequency, and channel conditions of the client. Using the historical data of the heterogeneous system running, based on the device conditions of each client, the optimal split point for each client to split the global model is determined with the goal of minimizing the total training delay. AsFigure 1 As shown in Figure 1 , the training method includes the following steps S101 to S107:
[0054] Step S101: During each round of training, all clients execute the training of the client local model in parallel. The client local model deployed on each client performs forward calculation using the training sample set constructed from the local data of each client, and outputs the shredded data.
[0055] Step S102: Each client transmits the shredded data and the labels of the training sample set to the server local model of the central server for forward calculation, and outputs the predicted values. Based on the predicted values and the labels, the loss value corresponding to each client is calculated.
[0056] Step S103: Calculate the gradients of the output layer of each server local model according to the loss values, perform backpropagation of the gradients in the server local model, obtain the gradients of the shredded data, and send them back to the corresponding clients.
[0057] Step S104: The client local model of each client performs backpropagation of the gradients of the shredded data to update the parameters of the client local model.
[0058] Step S105: After all client local models are updated, transmit all the updated client local models to the edge server.
[0059] Step S106: Set the network layers in the global model that contain the corresponding split point positions of all clients in the global model as the common layers. The edge server aggregates the parameters of the common layers involved in all the updated client local models, and the central server aggregates the parameters of the common layers involved in all the server local models after backpropagation. Exchange the information of the parameters of the common layers aggregated on both sides. The edge server aggregates all the client local models after information exchange to obtain the global client model. The central server aggregates all the server local models after information exchange to obtain the global server model.
[0060] Step S107: The edge server distributes the global client model to the corresponding clients according to the split points at which each client splits the global model.
[0061] Among them, the best split points at which each client splits the global model are determined with the goal of minimizing the total training latency. During the solution process, when the best split points solved by a client do not satisfy the feasible solutions, this client is not selected during training.
[0062] Among them, the shredded data is the intermediate result or feature representation generated after the local original data of the training sample set is processed by the client local model. This kind of processing helps to protect privacy while retaining the useful information for model training.
[0063] Among them, the global client model is equivalent to the set of client local models, and the global server model is equivalent to the set of server local models. They both include a common layer. The mutually matching client local model and server local model are combined into a global model.
[0064] In some embodiments of the present invention, the method further includes:
[0065] In each round of training, the training sample set of each client is divided into multiple mini-batch data, and the mini-batch gradient descent algorithm is used to iteratively update the parameters of this round of training.
[0066] In some embodiments of the present invention, for the batch size setting of the mini-batch data:
[0067] Set the minimum batch size and maximum batch size of the mini-batch data. Each time the training sample set is divided during training, the batch size is dynamically adjusted between the minimum batch size and the maximum batch size according to the current computing power and channel conditions of each client. The batch size of the mini-batch data is positively correlated with the computing power and the quality of the channel conditions.
[0068] For example, obtain the quantization values of the current computing power and channel conditions of each client in this round of training, and use linear interpolation or weighted average algorithm to calculate the batch size within the range of the minimum batch size and the maximum batch size.
[0069] In some embodiments of the present invention, to determine the best splitting point for each client to split the global model with the goal of minimizing the total training delay, it is based on optimizing the single-round training delay of each client. The single-round training delay of each client is:
[0070]
[0071] Among them, represents the forward propagation delay of the client local model, represents the forward propagation computation volume, f k represents the computing frequency of the client, κ k represents the computing density of the client;
[0072] represents the shredded data transmission delay, bξ s (l k ) represents the shredded data size, represents the uplink transmission rate from the client to the central server;
[0073] represents the forward propagation delay of the model on the central server, represents the computing frequency allocated by the central server to client k, κs Represents the computing density of the central server;
[0074] Represents the model backpropagation delay on the central server;
[0075] Represents the gradient backpropagation delay, bξ g (l k ) Represents the size of the shredded data gradient;
[0076] Represents the client local model backpropagation delay;
[0077] Represents the delay for the client local model to be transmitted to the edge server, ξ m (l k ) Represents the size of the client local model, Represents the uplink transmission rate from the client to the edge server;
[0078] Represents the delay for the edge server to distribute the client local model;
[0079] N represents the number of mini - batch data into which the training sample set is divided, obtained by dividing the training sample set by the batch size of the mini - batch data.
[0080] Among them, F represents forward meaning forward, B represents backward meaning backward, U represents upload meaning uplink, D represents download meaning downlink, ξ m In m represents model meaning model, ξ g In g represents gradients meaning gradients, without other special meanings.
[0081] In some embodiments of the present invention, before each round of training, the method further includes:
[0082] Based on each client local model determined by the optimal splitting point, with the goal of minimizing the total training delay, determine the optimal computing frequency assigned by the central server to each client, the optimal uplink and downlink transmission powers from each client to the central server, and the optimal uplink and downlink transmission powers from each client to the edge server.
[0083] In some embodiments of the present invention, the method for the edge server to aggregate the global client model and the central server to aggregate the global server model is: weighted average each layer of model parameters with the proportion of the local data volume of each client in the total data volume of all clients as the weight to obtain the aggregated model parameters.
[0084] In some embodiments of the present invention, the method further includes:
[0085] In a single training round, the edge server periodically sends signals to each client to confirm the online status of each client. If no response message is received within a preset time, the client is marked as an abnormal status.
[0086] For the clients marked as abnormal status in the current training round, the edge server abandons waiting for the clients to transmit the updated client local models, and only aggregates all the client local models except those in the abnormal status.
[0087] In some embodiments of the present invention, encrypted data transmission is performed between the client, the edge server, and the central server.
[0088] Further, symmetric key or asymmetric key encryption can be used for encrypted data transmission. A hash function can also be used to verify the transmitted data to detect whether the data has been tampered with during transmission.
[0089] Another embodiment of the present invention provides a method for training a split federated learning model based on a heterogeneous system, as Figure 2 shown. The implementation of this training method is based on the HSFL (Heterogeneous Split Federated Learning) framework, which includes a batch of users k ∈ {1, 2, …, K} of the client, a main server (MainServer, MS), and an edge server (Edge Server, ES), where the main server is the central server. The MS is responsible for the calculation of all server-side models, and the ES is responsible for the aggregation and distribution of the client models. Each user has a private dataset where: x k,i represents the original data, y k,i represents the corresponding label, |D k | represents the size of the dataset, and the data set of all users is To enable all resource-constrained users to participate in the training, each user k ∈ K needs to have an appropriate cut layer l k to make full use of its device resources and prevent falling behind during training.
[0090] By splitting the global model W into a client model k deployed on the user device and a server-side model deployed on the MS at layer l Let all users cooperate with the MS to complete parallel model training. When all users have completed the training task, the ES and MS will jointly complete model aggregation for the next round of training. Specifically, the workflow of HSFL can be divided into three stages. Let the training task require A training rounds, and the set of all training rounds is A. In training round a, the workflow is as follows:
[0091] 1. Forward propagation stage:
[0092] In this stage, users perform forward propagation in parallel using their private datasets, which specifically consists of the following three steps:
[0093] 1-i Client forward calculation: HSFL adopts the mini-batch gradient descent method and uses a mini-batch of data B k ={X k ,y k} for iteration. Each user performs the client model in parallel using the mini-batch of data to obtain the model output S k , that is where f() represents the network function of the model.
[0094] 1-ii Shredded data transmission: After completing the client forward calculation, the client transmits S k and the corresponding label y k to the MS.
[0095] 1-iii Server-side forward calculation: The MS uses S k as the input of the server-side model to complete the remaining part of the forward propagation and obtain the predicted value of the model Then, based on this, combined with the label y k calculate the average loss of this mini-batch where b is the size of the mini-batch and H() represents the cross-entropy loss function.
[0096] 2. Backward propagation stage:
[0097] In this stage, the MS and users use the calculated loss to update the model parameters. Similar to the forward propagation, there are also three steps:
[0098] 2-i Server-side backward calculation: Given the loss l(B k ; W k,a ), calculate the gradient The MS updates the server-side model with η as the learning rate
[0099] 2-ii Gradient transmission: After the server-side model completes the backward propagation, it obtains the gradient of the shredded data S k and transmits it back to the corresponding client.
[0100] 2-iii Client-side Reverse Calculation: After receiving the feedback backpropagation gradients, the client updates its local model.
[0101] The above forward propagation and backward propagation will be repeated N times until all mini-batches of the dataset are exhausted to ensure full utilization of the local dataset.
[0102] 3. Synchronous Aggregation Phase:
[0103] In the HSFL framework, different clients can choose their respective cut layers, resulting in differences in model parameters among users. Synchronous aggregation exchanges the common layer parameters through additional communication steps, aggregates the local updates of different clients, and generates a new global model to ensure that the updates of the aggregated global model on the common layer are consistent regardless of the cut layer selected by the client.
[0104] In this phase, the ES and MS cooperate to aggregate these updated models to form new global client-side and server-side models as the initial models for the next round of training. This phase is also divided into three steps:
[0105] 3-i Client Model Upload: After all users complete their respective training, they upload the updated models to the ES. To achieve synchronous aggregation, this step may cause waiting idle among users.
[0106] 3-ii Federated Aggregation: The ES and MS are responsible for federated aggregation of all client models and server-side models respectively to obtain the global model for the next round of training. Due to different cut layers, the client models and server-side models of different users are heterogeneous and cannot directly adopt the federated aggregation FedAvg method. Therefore, use L com to represent the common layer of the client model and the server-side model. The ES and MS need to exchange the parameters of their respective L com to complete the heterogeneous models into homogeneous models, and then perform weighted averaging according to the proportion of the local data volume of the users as the weights to obtain the new round of global client-side and server-side models.
[0107] 3-iii Model Distribution: After aggregation, the ES distributes the corresponding client models according to the cut layers of the users. This round of training ends.
[0108] As Figure 3 shown, the training latency composition of HSFL (Heterogeneous Split Federated Learning) includes:
[0109] 1-i: The forward propagation latency of the client model is:
[0110] Among them, represents the amount of computation required for forward propagation (FLOPs), which is related to the cutting layer l k is related. f k , κ k respectively represent the computing frequency and computing density of the user.
[0111] 1-ii: The latency of shredded data transmission is:
[0112] Among them, bξ s (l k ) represents the size of shredded data (bits), represents the uplink transmission rate from the user to the MS.
[0113] 1-iii: The latency of forward propagation of the server-side model is:
[0114] Among them, represents the computing frequency allocated by the MS to user k, κ s represents the computing density of the MS.
[0115] 2-i: The latency of backpropagation of the server-side model is:
[0116] 2-ii: The latency of gradient backpropagation is:
[0117] Among them, bξ g (l k ) represents the size of the gradient of shredded data.
[0118] 2-iii: The latency of backpropagation of the client-side model is:
[0119] 3-i: The upload latency of the client-side model
[0120] Among them, ξ m (l k ) represents the size of the client-side model (bits), represents the uplink transmission rate from the user to the ES.
[0121] 3-ii: Since the computation for aggregation only involves simple operations such as addition and averaging, its latency is ignored. At the same time, it is considered that the communication between the ES and the MS can be efficiently carried out through the power grid, so the latency of Lp transmission can also be ignored.
[0122] 3-iii: The latency of distributing the client-side model is:
[0123] Among them, F represents forward, B represents backward, U represents upload, D represents download, and ξ m in which m represents model, and ξ g in which g represents gradients, without other special meanings.
[0124] In summary, the total time delay of HSFL in each training round is:
[0125]
[0126] Therefore, the total time delay of a training task in the HSFL framework is According to the aforementioned time delay analysis of the training task, at the computing level, the cutting layer and computing power will directly affect the computing load and computing speed respectively, and thus affect the training delay. At the transmission level, the number and transmission power of the wireless channels allocated to the customers also determine the transmission time delay of the intermediate data (shredded data and gradients).
[0127] Therefore, the total time delay can be expressed as a function of the following variables:
[0128] Cutting layer allocation: l = {l1, l2,..., l K};
[0129] MS computing frequency allocation:
[0130] Communication resource allocation: where each element represents the up / down transmission power allocated to the user to the MS or ES.
[0131] Therefore, the objective function based on the total time delay can be obtained:
[0132]
[0133] Among them, C1 restricts the selection range of the cutting layer; C2 and C3 indicate that the MS computing frequency allocation cannot exceed the maximum computing frequency of the MS; C4 and C5 indicate the limitations of the transmission power.
[0134] For this NP-hard optimization problem, the user's storage resources are basically unchanged during one training session, indicating that the maximum model size that the customer can store is determined. Therefore, the cutting layer l can be fixed at the beginning of the training. Correspondingly, the user's communication ability will change with the variation of the channel quality during the training process. Considering that in a wireless environment, it is generally a quasi-static flat fading channel, that is, the channel gain remains basically unchanged within a short period of time but will change on a longer time scale. At the same time, the computing frequency allocation F of the MS will also change with the communication resources P. Therefore, the original problem can be decomposed into two sub-problems P L and P S , to optimize l and F, P respectively. The calling relationship of the resource optimization algorithm is as Figure 4 shown.
[0135] The specific step algorithm to solve this problem is as shown in the flowchart: Before the training starts, solve the long-time scale problem to obtain the optimal cutting layer l * After that, based on this, continue to solve the short-time scale problem in each round of training.
[0136] Among them, the long-time scale problem is described as follows:
[0137]
[0138] For this problem, one of the difficulties lies in how to reasonably handle the dynamically changing random variables, the user equipment computing frequency f and the channel condition G, and estimate the total training delay together with the optimization variables. In addition, HSFL allows users to select different cutting layers, resulting in a huge solution space for the combination of cutting layers. To solve these problems, this embodiment adopts a genetic algorithm (GA) based on the sample average approximation (SAA) idea to solve the optimal cutting layer selection in the system. Among them, SAA uses the delay of each round calculated by the historical data samples of f and G to approximate the average delay per round T a . GA can more efficiently find the cutting layer selection l that makes the smallest. Therefore, the objective function of P L can be expressed as
[0139]
[0140] The solution steps of this algorithm are as follows:
[0141] First, assume that the data of f and G satisfy the normal distribution with their respective means and variances. Then, the algorithm in P S can be used to solve F and P when l, f and G are known. Then, use s (sufficiently large) samples of f and G to calculate the single-round delay respectively and use their mean value to approximate the expected value Specifically expressed as Among them, T(l,F s ,P s ; G s ,f s ) represents the optimal one-round delay under the sample s. Finally, use GA to find l * .
[0142] Specifically, based on the standard GA, this algorithm changes the fitness function to Since GA generally optimizes in the direction of higher fitness of individuals, a negative sign needs to be added before the function to represent minimizing the delay. More specifically, initialize a population containing P individuals (representing the selection of the cutting layer), where the fitness of individual l p is expressed as Then iterate several times to find the best individual. When the fitness of the best individual does not change significantly for g generations, it is considered that l * is found.
[0143] For short-time scale problems P S , where F and P are still coupled, and the joint solution is still very difficult. Therefore, it can be decomposed into P S-F , P S-P to optimize F and P respectively, that is:
[0144]
[0145] For P S-F , in the one-round delay of the user, only the calculation time on the MS is related to F. When the cutting layer and the corresponding P, f, G are given, for this constrained univariate problem, auxiliary variables can be introduced to simplify the problem and the Lagrangian relaxation method can be used to solve it. Specifically, P S-F can be expressed as:
[0146]
[0147] Among them, represents the value unrelated to F in the one-round delay of user k and can be directly calculated, and can be regarded as a constant.
[0148] For solving this max-min problem, first introduce C6: to simplify the problem to:
[0149]
[0150] Then construct the Lagrangian function:
[0151]
[0152] Given the initial F, λ, μ k , the initial objective function can be directly calculated. By taking the partial derivatives of L with respect to each variable and iteratively updating F, λ, μ k , the entire process will end at , where σ is a preset small value.
[0153] For P S-P , in HSFL, the client does not communicate with the MS and ES simultaneously and uses different frequency bands. Therefore, the communication resources from the user to the MS and from the user to the ES can be optimized separately in the same way, that is Therefore, P S-P can be further divided into two parts:
[0154]
[0155] To solve such problems, the branch and bound method is usually adopted. This method decouples the problem and calculates the lower bounds of sub-problems to eliminate unreachable solutions and then gradually approaches the optimal solution. In this embodiment, the pysciopt solver with the branch and bound method is used to solve and
[0156] In summary, the present invention discloses a method for training a split federated learning model based on a heterogeneous system. The system includes multiple clients, a central server, and an edge server. The global model divides the client local model and the server local model based on each client device condition and deploys them to the corresponding client and central server respectively. In each round of training, the clients train in parallel, forward-propagate the local data to output shredded data and transmit it to the server local model of the central server to continue forward-propagation to obtain a loss value, and then back-propagate the gradient to update the parameters of the server local model and the client local model. Each client transmits the updated model to the edge server. After the edge server and the central server complete the parameter exchange of the common layer, they perform a model aggregation to generate a global client model and a global server model. The global client model distributes the client local model to the clients according to the split points of each client. The present invention adapts to the conditions of client devices to perform personalized segmentation of the client model, reduces the computational burden of heterogeneous devices, and improves the model training efficiency of heterogeneous systems.
[0157] Correspondingly to the above method, the present invention further provides a device / system, which includes a computer device. The computer device includes a processor and a memory. Computer instructions are stored in the memory, and the processor is configured to execute the computer instructions stored in the memory. When the computer instructions are executed by the processor, the device / system implements the steps of the method as described above.
[0158] An embodiment of the present invention further provides a computer-readable storage medium, on which a computer program is stored. When the computer program is executed by a processor, it implements the steps of the aforementioned edge computing server deployment method. The computer-readable storage medium may be a tangible storage medium, such as a random access memory (RAM), internal memory, read-only memory (ROM), electrically programmable ROM, electrically erasable programmable ROM, register, floppy disk, hard disk, removable storage disk, CD-ROM, or any other form of storage medium well-known in the technical field.
[0159] Those of ordinary skill in the art should understand that the various exemplary components, systems, and methods described in connection with the embodiments disclosed herein can be implemented in hardware, software, or a combination of both. Specifically, whether to implement in hardware or software depends on the specific application and design constraints of the technical solution. A professional technician can use different methods to implement the described functions for each specific application, but such implementation should not be considered to exceed the scope of the present invention. When implemented in hardware, it can be, for example, an electronic circuit, an application-specific integrated circuit (ASIC), appropriate firmware, a plug-in, a functional card, and so on. When implemented in software, the elements of the present invention are programs or code segments used to perform the required tasks. The program or code segment can be stored in a machine-readable medium or transmitted through a data signal carried in a carrier wave on a transmission medium or a communication link.
[0160] It should be clear that the present invention is not limited to the specific configurations and processes described above and shown in the figures. For the sake of brevity, the detailed description of known methods is omitted here. In the above embodiments, several specific steps are described and shown as examples. However, the method process of the present invention is not limited to the specific steps described and shown. Those skilled in the art can make various changes, modifications, and additions, or change the order between steps after understanding the spirit of the present invention.
[0161] In the present invention, the features described and / or illustrated for one embodiment can be used in the same or a similar manner in one or more other embodiments, and / or combined with the features of other embodiments or replace the features of other embodiments.
[0162] The above are only the preferred embodiments of the present invention and are not intended to limit the present invention. For those skilled in the art, various modifications and variations can be made to the embodiments of the present invention. Any modification, equivalent replacement, improvement, etc. made within the spirit and principle of the present invention shall be included within the protection scope of the present invention.
Claims
1. A method for training a segmented federated learning model based on a heterogeneous system, characterized in that: The heterogeneous system includes multiple clients, a central server and an edge server; The global model is divided into a client local model adapted to each client and a corresponding server local model based on the device condition of each client; the client local model is used to be deployed on the corresponding client, and the server local model is used to be deployed on the central server; wherein the device condition of each client includes the device storage resources, computing frequency and channel condition of the client, and the optimal segmentation point for each client to segment the global model is determined based on the device condition of each client with the goal of minimizing the total training delay by using the historical data of the operation of the heterogeneous system; the training method comprises the following steps: In each round of training, all clients execute the training of the client local model in parallel, and the client local model deployed on each client performs forward calculation using the training sample set constructed by each client local data, and outputs crushed data; Each client transmits the crushed data and the label of the training sample set to the server local model of the central server for forward calculation, and outputs a predicted value; based on the predicted value and the label, the loss value corresponding to each client is calculated; Calculate the gradient of the output layer of each server local model according to the loss value, back-propagate the gradient in the server local model, obtain the gradient of the crushed data and transmit it back to the corresponding client; The client local model of each client back-propagates the gradient of the crushed data to update the parameters of the client local model; After all client local models are updated, all updated client local models are transmitted to the edge server; The network layer in the global model that contains the corresponding segmentation point position of all clients in the global model is set as a common layer; the edge server collects the parameters related to the common layer in all updated client local models, and the central server collects the parameters related to the common layer in all server local models after back propagation, and the parameters of the common layer collected by both sides are exchanged; the edge server aggregates all client local models after information exchange to obtain a global client model; the central server aggregates all server local models after information exchange to obtain a global server model; The edge server divides the global client model into division points of the global model for each client, and distributes the client local model to the corresponding client.
2. The method for training a segmented federated learning model based on a heterogeneous system according to claim 1, characterized in that: The method further includes: In each round of training, the training sample set of each client is divided into multiple small batches of data, and the small batch gradient descent algorithm is used to iteratively update the parameters of this round of training.
3. The method for training a segmented federated learning model based on a heterogeneous system according to claim 2 is characterized in that: For the mini-batch data the batch size is set as: Set the minimum batch size and maximum batch size of small batch data; when dividing the training sample set each time for training, dynamically adjust the batch size between the minimum batch size and the maximum batch size according to the current computing power and channel conditions of each client; the batch size of small batch data is positively correlated with the computing power and the quality of channel conditions.
4. The method for training a segmented federated learning model based on a heterogeneous system according to claim 2, characterized in that: The optimal split point for each client to split the global model is determined with the goal of minimizing the total training delay, which is based on optimizing the single-round training delay of each client. The single-round training delay of each client is: Among them, l k Indicates the cutting layer; k indicates the client number, K indicates the client set, and max indicates the maximum value; represents the forward propagation delay of the client local model, represents the amount of forward propagation calculation, f k represents the computation frequency of the client, κ k Indicates the computing density of the client; represents the pulverized data transmission delay, bξ s (l k ) indicates the size of the crushed data, Indicates the uplink transmission rate from the client to the central server; represents the forward propagation delay of the model on the central server, represents the computation frequency assigned by the central server to client k, κ s Indicates the computing density of the central server; Represents the model back propagation delay on the central server; Denotes the gradient return delay, bξ g (l k ) represents the magnitude of the gradient of the crushed data; g represents the gradient; represents the back propagation delay of the client local model; represents the delay of transmitting the local model from the client to the edge server, ξ m (l k ) represents the client local model size, represents the uplink transmission rate from the client to the edge server; m represents the model; represents the time delay of the edge server distributing the client local model; N represents the number of small batches of data into which the training sample set is divided, which is obtained by dividing the training sample set by the batch size of the small batch data; F represents forward, B represents reverse, U represents uplink, D represents downlink, and b represents the size of the batch; MS represents the central server, and ES represents the edge server.
5. The method for training a segmented federated learning model based on a heterogeneous system according to claim 1, characterized in that: Before each round of training, the method also includes: Based on the local model of each client determined by the optimal split point, with the goal of minimizing the total training delay, the optimal computing frequency assigned to each client by the central server, the optimal uplink and downlink transmission power from each client to the central server, and the optimal uplink and downlink transmission power from each client to the edge server are determined.
6. The method for training a segmented federated learning model based on a heterogeneous system according to claim 1, characterized in that: The edge servers aggregate to obtain a global client model, and the central servers aggregate to obtain a global server model by weighting the parameters of each layer of the model using the proportion of each client's local data volume to the total data volume of all clients as a weight to obtain the aggregated model parameters.
7. The method for training a segmented federated learning model based on a heterogeneous system according to claim 1, characterized in that: The method further includes: In a single training round, the edge server periodically sends a signal to each client to confirm the online status of each client. If no reply information is received within a preset time, the client is marked as abnormal. For a client marked as being in an abnormal state in the current training round, the edge server gives up waiting for the client to transmit an updated client local model, and only aggregates all client local models except those in the abnormal state.
8. The method for training a segmented federated learning model based on a heterogeneous system according to claim 1, characterized in that: Encrypted data transmission is performed between the client and the edge server and the central server.
9. A computer-readable storage medium having a computer program / instruction stored thereon, characterized in that: When the computer program / instructions are executed by a processor, the steps of the method as claimed in any one of claims 1 to 7 are implemented.
10. A computer program product comprising a computer program / instructions, characterized in that When the computer program / instructions are executed by a processor, the steps of the method according to any one of claims 1 to 7 are implemented.
Citation Information
Patent Citations
Federal learning training method and system based on model segmentation and resource allocation
CN114925852A
Efficient communication federal learning method for semantic segmentation of small sample medical images
CN115965782A