A traversal federated learning method and system based on tensor filling
By adopting a traversal federated learning method based on tensor padding in federated learning, the data accuracy and computing efficiency problems caused by device heterogeneity are solved, and the uniform training of the global model in low-capacity devices is achieved.
Patent Information
- Application Number
- CN202510217551.1
- Authority / Receiving Office
- CN · China
- Patent Type
- Patents(China)
- Current Assignee / Owner
- Filing Date
- 2025-02-26
- Publication Date
- 2025-05-06
- Estimated Expiration
- 2045-02-26
AI Technical Summary
The heterogeneous federated learning of devices affects the accuracy and computational efficiency of data due to the differences in hardware specifications between different devices, or cannot accommodate a complete model for training.
Using a traversal federated learning method based on tensor padding, the bit width recovery component and model segmentation component are deployed on the server side, and the gradient is hierarchically modeled based on the device bit width, and different bit width gradients are aligned through tensor padding technology to weight aggregate to generate a new global model. At the same time, the traversal window mechanism is used to segment the model, so that the global model is uniformly trained in low-capacity devices.
Effectively coordinate and integrate model updates of different devices, improve data accuracy and computing efficiency, ensure that the resources of all edge devices can be maximized, and reduce resource waste and energy consumption.
Smart Images

Figure CN119721187B_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the technical field of Internet of Things and relates to a traversal federated learning method and system based on tensor filling. Background Art
[0002] With the development of artificial intelligence Internet of Things (AIoT) systems, the hardware technology of edge devices has been greatly improved, and the data processing and communication capabilities have been continuously enhanced. In this context, in order to solve the problem of data silos and protect data privacy, federated learning has emerged as a cutting-edge distributed machine learning strategy. It avoids the risk of storing and transmitting sensitive data to central servers, while making full use of the computing power of edge devices for model training. It has applications in fields such as healthcare and industrial engineering. However, it is worth noting that with the continuous expansion of the scale of neural network models, the demand for computing resources in the federated learning framework has also risen sharply, while the computing and storage resources of some edge devices are relatively limited, which to some extent limits their ability to process large-scale data and complex computing tasks. Many edge devices in the AIoT ecosystem have a negative impact on the accuracy of the global model due to their low bit width and limited storage capacity. Traditional federated learning frameworks usually exclude these edge devices, resulting in a large amount of resource waste. In the context of increasing emphasis on low-carbon sustainable computing, it is particularly important to maximize the use of all available edge devices to reduce resource waste and energy consumption. To solve this problem, the concept of device-heterogeneous federated learning was proposed and further developed. It breaks through the boundaries of traditional federated learning and allows different types of edge devices (such as smart watches, sensors, etc.) to participate in the federated learning process. Each device performs model training according to its specific functions and resource limitations. However, device heterogeneity also brings significant challenges. The hardware specifications of different devices may vary greatly. Some devices may build models based on low-bitwidth hardware operations (such as using FPGA, ASIC, Raspberry Pi, or Edge GPUs), which directly affects the accuracy and computing efficiency of the data. At the same time, there are also some devices that cannot accommodate the complete model for training. How to effectively coordinate and integrate these model updates from different devices has become a key problem. Summary of the invention
[0003] The purpose of the present invention is to solve the problem in the prior art that heterogeneous federated learning of devices affects data accuracy and computing efficiency due to differences in hardware specifications between different devices, or cannot accommodate a complete model for training, and to provide a traversal federated learning method and system based on tensor filling.
[0004] In order to achieve the above object, the present invention adopts the following technical solutions:
[0005] A traversal federated learning method based on tensor filling, comprising:
[0006] Deploy the bit width recovery component and the model segmentation component on the server side;
[0007] Perform layered tensor modeling on the corresponding gradients based on the device bit width; classify the parameters of the same model layer in different gradients on the server side, and model each layer as a tensor;
[0008] Based on the tensor filling technology, the different bit width gradients corresponding to the tensors in the bit width recovery component are aligned, and all gradients are weightedly aggregated, and a new global model is generated based on the aggregated weights;
[0009] Based on the model segmentation component, a traversal window mechanism is constructed to segment the model. Through the traversal window mechanism, different layers of the global model are selected in the server and the client for each round of training, so that the global model can be evenly trained in low-capacity devices until the global model converges.
[0010] A further improvement of the present invention is:
[0011] Furthermore, the corresponding gradients are modeled as hierarchical tensors based on the device bit width, specifically:
[0012] Identify model layers of the neural network that require bit-width-based tensor modeling;
[0013] Based on the determined model layer, identify different data bit widths supported by the devices participating in the training, and determine the bit width dimension of the tensor accordingly;
[0014] Map the model parameters trained by each device at this level to the corresponding position in the tensor according to its bit width; aggregate all device parameters under the same bit width to form a sub-tensor or element set;
[0015] Parameter bit width mapping is performed independently for each layer of the neural network to ensure that the tensor can accurately reflect the parameter differences between different bit width devices on that layer.
[0016] Furthermore, based on the tensor filling technology, the different bit width gradients corresponding to the tensors in the bit width recovery component are aligned. Specifically, the tensor size of the highest bit width device is used as the standard tensor, and the tensor of the low bit width device is inserted with zero values to reach the standard tensor; the missing values in the low bit width data are filled by tensor decomposition and tensor reconstruction, and the tensor of the low bit width model parameters is aligned with the tensor shape of the full precision model parameters.
[0017] Furthermore, tensor decomposition includes Tucker decomposition and CP decomposition; the Tucker decomposition is to convert a N The tensor is decomposed into a N The product of a core tensor of order and multiple factor matrices; for a NA tensor of order, of the form:
[0018] (1);
[0019] Among them, the elements in the tensor are represented as , is the element index, is the real number space, is the size of the tensor in each dimension. For example, for a three-dimensional tensor with a size of 10×20×30, N=3. That is, there are 10 elements in the first dimension, 20 elements in the second dimension, and 30 elements in the third dimension. The tensor elements are represented as , and 1≤i1≤10, 1≤i2≤20, 1≤i3≤30, the tensor element is expressed as ; , and 1≤i1≤10, 1≤i2≤20, 1≤i3≤30;
[0020] With a third-order tensor For example, its tucker decomposition is as follows:
[0021] (2)
[0022] in, represents the core tensor, is the dimension size of the core tensor, A, B, C represent factor matrices, represents the outer product; , and The factor matrices are A , B and C Column vector in ; is the core tensor G Elements in
[0023] The CP decomposition is a method of decomposing the target tensor into the sum of multiple rank 1 tensors. N Rank Tensor , its CP is decomposed into the following form:
[0024] (3);
[0025] Among them, r is the rank of the decomposition selection, represents the outer product, Represents a factor matrix, each factor matrix belong , is the real number space, It represents the size of the original tensor in the i-th dimension. For example, for a three-dimensional tensor with a size of 10×20×30, the size of the first dimension is 10, that is, n1=10, the size of the second dimension is 20, that is, n2=20, and the size of the third dimension is 30, that is, n3=10. Select r=3 and substitute its cp decomposition into formula (3) to obtain , is a column vector in the first factor matrix, belonging to , and so on, and Belong to and , when t=1, , , They are , , The first column of The first rank 1 tensor is obtained, and so on. When t takes the value of 2 and 3, the second and third rank 1 tensors can be obtained respectively, thus completing the decomposition process. The tensor reconstruction is the inverse calculation of tensor decomposition, which is used to restore the tensor shape and approximate the original tensor; the CP decomposition attempts to find the minimum rank so that Approximate the original tensor as accurately as possible.
[0026] Furthermore, the shape of the tensor of the low-bit-width model parameters is aligned with the tensor shape of the full-precision model parameters. Specifically, the low-bit-width model parameter elements are converted into multiple matrices and expanded by adding column vectors and filling zero elements; the expanded matrices are decomposed and reconstructed using CP decomposition or Tucker decomposition; after multiple rounds of iterations, a completely filled tensor is finally obtained; the newly added columns are merged into the original tensor and the shape is restored to complete the dequantization process of the model parameters.
[0027] Furthermore, all gradients are weighted aggregated, and a new global model is generated based on the aggregated weights. Specifically, a weight is assigned to the gradient of each device or client based on the original bit width of the parameter; the weight is usually proportional to the accuracy of the bit width; and the gradients from different devices or clients are weighted aggregated using the determined weights.
[0028] Furthermore, based on the model segmentation component, a traversal window mechanism is constructed to segment the model. Through the traversal window mechanism, different layers of the global model are selected in the server and the client for each round of training, so that the global model can be evenly trained in low-capacity devices until the global model converges. Specifically, in each round of training in federated learning, the server extracts sub-models of different capacities from the global model using the traversal window, and broadcasts them to clients with corresponding computing capabilities respectively; the client trains the received sub-models according to local data, and transmits the updated gradients of their heterogeneous sub-models to the server; the server summarizes the updated gradients and uses them to update the global model for the next round; the traversal window advances in each round, and loops through all parts of the global model in turn in different rounds; this process is iterated continuously until the global model reaches a state of uniform training after multiple trainings until it converges.
[0029] Furthermore, the server uses the traversal window to extract sub-models of different capacities from the global model and broadcasts them to clients with corresponding computing capabilities, specifically:
[0030] Assuming a global model have L Layer, j The sub-models extracted by the server in round 1 include i Layer to The parameters of the layer; the traversal window moves forward one step in each round, skipping the model layer taken in the previous round; assuming that in the j In the round, the first n The client model is , then the traversal window starts from the global model Extracted from i Layer to The method of the layer sub-model is expressed as:
[0031] (4);
[0032] in, Indicated in j In round, n Each client extracts the starting layer index of the sub-model; k is the size of the sub-model, that is, the number of layers of the sub-model. The traversal window will advance 1 layer per round during the training process and L The update position of the traversal window is expressed as:
[0033] (5).
[0034] Furthermore, the client trains the received sub-models based on local data and transmits the updated gradients of their heterogeneous sub-models to the server. Specifically, since the local client only trains different sub-models in the global model in each round, the aggregation object is the heterogeneous sub-model extracted by the traversal window. For this heterogeneous model, the parameter part trained in this round is weighted and aggregated according to the aggregation strategy. The aggregation process is expressed as:
[0035] (6);
[0036] in, is the global model after aggregation, represents the learning rate, Indicated in j In round, n Client-to-submodel The updated gradient generated after training, N Represents the number of clients participating in the training, Representative n The local data of each client.
[0037] A traversal federated learning system based on tensor filling, comprising:
[0038] A deployment module, wherein the deployment module deploys the bit width recovery component and the model segmentation component on the server side;
[0039] A tensor modeling module, wherein the tensor modeling module performs hierarchical tensor modeling on corresponding gradients based on device bit width; classifying parameters of the same model layer in different gradients on the server side, and modeling each layer as a tensor;
[0040] An alignment module, wherein the alignment module aligns different bit width gradients corresponding to tensors in the bit width recovery component based on tensor filling technology, performs weighted aggregation on all gradients, and generates a new global model based on the aggregated weights;
[0041] The traversal module constructs a traversal window mechanism to segment the model based on the model segmentation component. Through the traversal window mechanism, different layers of the global model are selected in the server and the client for each round of training, so that the global model is evenly trained in the low-capacity device until the global model converges. BRIEF DESCRIPTION OF THE DRAWINGS
[0042] In order to more clearly illustrate the technical solutions of the embodiments of the present invention, the drawings required for use in the embodiments are briefly introduced below. It should be understood that the following drawings only show certain embodiments of the present invention and therefore should not be regarded as limiting the scope. For ordinary technicians in this field, other related drawings can be obtained based on these drawings without creative work.
[0043] Figure 1 It is a flowchart of the traversal federated learning method based on tensor filling of the present invention;
[0044] Figure 2 It is a structural diagram of the ergodic federated learning system based on tensor filling of the present invention;
[0045] Figure 3 This is a comparison diagram of the training loss of Resnet18 on the MNIST dataset at different learning rates;
[0046] Figure 4 This is a comparison diagram of the training loss of Resnet18 on the CIFAR-10 dataset at different learning rates;
[0047] Figure 5 This is a schematic diagram comparing the changes in global model accuracy of Resnet18 for cp decomposition and tucker decomposition on the MNIST dataset;
[0048] Figure 6 This is a schematic diagram comparing the changes in global model accuracy of Resnet18 after CP decomposition and Tucker decomposition on the CIFAR-10 dataset. DETAILED DESCRIPTION
[0049] In order to make the purpose, technical solutions and advantages of the embodiments of the present invention clearer, the technical solutions in the embodiments of the present invention will be clearly and completely described below in conjunction with the drawings in the embodiments of the present invention. Obviously, the described embodiments are part of the embodiments of the present invention, not all of the embodiments. Generally, the components of the embodiments of the present invention described and shown in the drawings here can be arranged and designed in various different configurations.
[0050] Therefore, the following detailed description of the embodiments of the present invention provided in the accompanying drawings is not intended to limit the scope of the invention claimed for protection, but merely represents selected embodiments of the present invention. Based on the embodiments of the present invention, all other embodiments obtained by ordinary technicians in this field without creative work are within the scope of protection of the present invention.
[0051] It should be noted that similar reference numerals and letters denote similar items in the following drawings, and therefore, once an item is defined in one drawing, further definition and explanation thereof is not required in subsequent drawings.
[0052] In the description of the embodiments of the present invention, it should be noted that if the terms "upper", "lower", "horizontal", "inner", etc. indicate an orientation or positional relationship based on the orientation or positional relationship shown in the drawings, or the orientation or positional relationship in which the product of the invention is usually placed when in use, it is only for the convenience of describing the present invention and simplifying the description, and does not indicate or imply that the device or element referred to must have a specific orientation, be constructed and operated in a specific orientation, and therefore cannot be understood as a limitation on the present invention. In addition, the terms "first", "second", etc. are only used to distinguish the description, and cannot be understood as indicating or implying relative importance.
[0053] In addition, if the term "horizontal" appears, it does not mean that the component must be absolutely horizontal, but can be slightly tilted. For example, "horizontal" only means that its direction is more horizontal than "vertical", which does not mean that the structure must be completely horizontal, but can be slightly tilted.
[0054] In the description of the embodiments of the present invention, it is also necessary to explain that, unless otherwise clearly specified and limited, the terms "set", "install", "connect", and "connect" should be understood in a broad sense, for example, it can be a fixed connection, a detachable connection, or an integral connection; it can be a mechanical connection or an electrical connection; it can be a direct connection, or it can be indirectly connected through an intermediate medium, or it can be the internal connection of two components. For ordinary technicians in this field, the specific meanings of the above terms in the present invention can be understood according to specific circumstances.
[0055] The present invention is further described in detail below in conjunction with the accompanying drawings:
[0056] See also Figure 1 The present invention discloses a traversal federated learning method based on tensor filling, comprising:
[0057] S101, deploying a bit width recovery component and a model segmentation component on a server side;
[0058] In the artificial intelligence IoT scenario constructed by the present invention, edge devices are composed of full-width high-capacity devices and low-width low-capacity devices. The server collects capacity and accuracy requirements from multiple clients. This information is used to match the distributed model and calculate the aggregation weight to ensure that the contribution of each client can be used fairly and effectively in the subsequent training process.
[0059] S102, performing layered tensor modeling on the corresponding gradients based on the device bit width; classifying the parameters of the same model layer in different gradients on the server side, and modeling each layer as a tensor;
[0060] The corresponding gradients are modeled as hierarchical tensors based on the device bit width, specifically:
[0061] Identify model layers of the neural network that require bit-width-based tensor modeling;
[0062] Based on the determined model layer, identify different data bit widths supported by the devices participating in the training, and determine the bit width dimension of the tensor accordingly;
[0063] Map the model parameters trained by each device at this level to the corresponding position in the tensor according to its bit width; aggregate all device parameters under the same bit width to form a sub-tensor or element set;
[0064] Parameter bit width mapping is performed independently for each layer of the neural network to ensure that the tensor can accurately reflect the parameter differences between different bit width devices on that layer.
[0065] S103, aligning different bit width gradients corresponding to tensors in the bit width recovery component based on the tensor filling technology, performing weighted aggregation on all gradients, and generating a new global model based on the aggregated weights;
[0066] Based on the tensor filling technology, the different bit width gradients corresponding to the tensors in the bit width recovery component are aligned. Specifically, the tensor size of the highest bit width device is used as the standard tensor, and the tensor of the low bit width device is inserted with zero values to reach the standard tensor; the missing values in the low bit width data are filled by tensor decomposition and tensor reconstruction, and the tensor of the low bit width model parameters is aligned with the tensor shape of the full precision model parameters.
[0067] Tensor decomposition includes Tucker decomposition and CP decomposition; the Tucker decomposition is to convert a N The tensor is decomposed into a N The product of a core tensor of order and multiple factor matrices; for a N A tensor of order, of the form:
[0068] (1);
[0069] Among them, the elements in the tensor are represented as ,in ;
[0070] With a third-order tensor For example, its tucker decomposition is as follows:
[0071] (2);
[0072] in, represents the core tensor, is the dimension size of the core tensor, A, B, C represent factor matrices, represents the outer product; , and The factor matrices are A , B and C Column vector in ; is the core tensor G Elements in
[0073] The CP decomposition is a method of decomposing the target tensor into the sum of multiple rank 1 tensors. N Rank Tensor , its CP is decomposed into the following form:
[0074] (3);
[0075] in, represents the outer product, Represents a factor matrix, each factor matrix belong , is the real number space, It represents the size of the original tensor in the i-th dimension, for example, a three-dimensional tensor of size 10×20×30, for the first dimension (i.e., i=1), n=10. i is a positive integer from 1 to N, and N is the rank of the tensor; the tensor reconstruction is the inverse calculation of tensor decomposition, which is used to restore the tensor shape and approximate the original tensor; the CP decomposition attempts to find the minimum rank R , so that formula (3) approximates the original tensor as accurately as possible.
[0076] The shape of the tensor of low-bit-width model parameters is aligned with the shape of the tensor of full-precision model parameters. Specifically, the low-bit-width model parameter elements are converted into multiple matrices, and expanded by adding column vectors and filling zero elements; the expanded matrices are decomposed and reconstructed using CP decomposition or Tucker decomposition; after multiple rounds of iterations, a fully filled tensor is finally obtained; the newly added columns are merged into the original tensor and the shape is restored to complete the dequantization process of the model parameters.
[0077] All gradients are weighted aggregated, and a new global model is generated based on the aggregated weights. Specifically, a weight is assigned to the gradient of each device or client based on the original bit width of the parameter; the weight is usually proportional to the accuracy of the bit width; and the gradients from different devices or clients are weighted aggregated using the determined weights.
[0078] S104, based on the model segmentation component, construct a traversal window mechanism segmentation model. Through the traversal window mechanism, different layers of the global model are selected in the server and the client for each round of training, so that the global model is evenly trained in the low-capacity device until the global model converges.
[0079] In each round of training in federated learning, the server uses the traversal window to extract sub-models of different capacities from the global model and broadcasts them to clients with corresponding computing capabilities respectively; the client trains the received sub-models based on local data and transmits the updated gradients of their heterogeneous sub-models to the server; the server summarizes the updated gradients and uses them to update the global model for the next round; the traversal window advances in each round and loops through all parts of the global model in turn in different rounds; this process is iterated continuously until the global model reaches a state of uniform training after multiple trainings until convergence.
[0080] The server uses the traversal window to extract sub-models of different capacities from the global model and broadcasts them to clients with corresponding computing capabilities, specifically:
[0081] Assuming a global model have L Layer, j The sub-models extracted by the server in round 1 include i Layer to The parameters of the layer; the traversal window moves forward one step in each round, skipping the model layer taken in the previous round; assuming that in the j In the round, the first n The client model is , then the traversal window starts from the global model Extract from i Layer to The method of the layer sub-model is expressed as:
[0082] (4);
[0083] in, Indicated in j In round, n Each client extracts the starting layer index of the sub-model; k is the size of the sub-model, that is, the number of layers of the sub-model. The traversal window will advance 1 layer per round during the training process and L The update position of the traversal window is expressed as:
[0084] (5);
[0085] The client trains the received sub-models based on local data and transmits the updated gradients of their heterogeneous sub-models to the server. Specifically, since the local client only trains different sub-models in the global model in each round, the aggregation object is the heterogeneous sub-model extracted by the traversal window. For this heterogeneous model, the parameter part trained in this round is weighted and aggregated according to the aggregation strategy. The aggregation process is expressed as:
[0086] (6);
[0087] in, is the global model after aggregation, represents the learning rate, Indicated in j In round, n Client-to-submodel The updated gradient generated after training, N Represents the number of clients participating in the training, Representative n The local data of each client.
[0088] The core advantage of this strategy is that it allows resource-constrained devices to store and process only smaller sub-models, thereby reducing the burden of storage and computing. At the same time, through the mechanism of traversing windows, it can ensure that the global model is trained evenly in multiple rounds of training, avoiding the problem of over-training or under-training of certain parameters. In addition, the aggregation strategy only performs weighted aggregation on the trained parameter part, thereby improving the accuracy and stability of the aggregation results.
[0089] Entering the training and updating phase, the algorithm adopts a parallelization strategy to allow each client to independently train the sub-model on its local data and update its parameters. Subsequently, these updated parameters are uploaded to the server, which is responsible for restoring the accuracy of the global model. If the accuracy of the model parameters does not meet the float32 standard, the algorithm will use tensor filling technology to restore the information gap between the low bit width and the high bit width, thereby improving the accuracy of the model. Next, the server uses the aggregation weights to perform weighted aggregation on all uploaded parameters to generate new global model parameters, which are used to update the global model. This process will be repeated until the preset maximum number of iterations is reached.
[0090] See also Figure 2 The present invention discloses a traversal federated learning system based on tensor filling, comprising:
[0091] A deployment module, wherein the deployment module deploys the bit width recovery component and the model segmentation component on the server side;
[0092] A tensor modeling module, wherein the tensor modeling module performs hierarchical tensor modeling on corresponding gradients based on device bit width; classifying parameters of the same model layer in different gradients on the server side, and modeling each layer as a tensor;
[0093] An alignment module, wherein the alignment module aligns different bit width gradients corresponding to tensors in the bit width recovery component based on tensor filling technology, performs weighted aggregation on all gradients, and generates a new global model based on the aggregated weights;
[0094] The traversal module constructs a traversal window mechanism to segment the model based on the model segmentation component. Through the traversal window mechanism, different layers of the global model are selected in the server and the client for each round of training, so that the global model is evenly trained in the low-capacity device until the global model converges.
[0095] Example:
[0096] Analysis of experimental results: In the artificial intelligence Internet of Things constructed in this example, the accuracy of the model under the federated learning framework it builds is the most intuitive verification indicator. The present invention simulates a FL system with 10 clients, in which the participating clients have different bit width configurations, capacity configurations and local data distributions, and the data distribution of each client is generated according to the Dirichlet distribution DirK(α). The pre-activated ResNet18 model is trained on CIFAR-10 and MNIST, static batch normalization is performed in PreResNet18, and a scalar module is added after each convolutional layer. The present invention deploys two bit width settings of 8bit and 32bit, and the capacity is evenly set to ,To generate partial models of corresponding sizes, the number of convolutional layer kernels of ResNet18 was modified, and the nodes in the output layer were kept the same.
[0097] Figure 3 and Figure 4 The global model loss values under various learning rates were compared to identify the optimal learning rate configuration that can maximize the model convergence efficiency and final performance. The experiment covers multiple predefined learning rate values, monitors and records the loss value of the model on an independent test set, and this process ensures the objectivity and repeatability of the experimental results. The results intuitively show the comparison of the loss values of the ResNet18 model under different learning rate configurations on the CIFAR-10 and MNIST datasets. Through careful comparison and evaluation, the learning rate of 0.0002 shows a significant advantage at the global loss level. Under this learning rate setting, the model performs better than other candidate values on both CIFAR-10 and MNIST benchmark datasets. Based on the above findings, in order to maintain the consistency of experimental conditions and optimize the efficiency and effect of subsequent research, the subsequent experiments will uniformly set the learning rate to 0.0002. This decision is intended to maximize the use of previous experimental results, promote further improvement of the overall performance of the model, and reduce potential biases introduced by improper hyperparameter selection.
[0098] Table 1 compares the performance of the present invention with fedrolex, fedrated dropout, and HeteroFL methods. At the same time, in order to simulate real low-bitwidth and low-capacity devices, two groups of experiments were set up, one group was 50% Int8 and 50% Float32 models, and the other group was 80% Int8 and 20% Float32 models. iid means independent and identically distributed, and Non-iid means non-independent and identically distributed. On the CIFAR-10 and MNIST data sets, whether it is a 50% Int8-50% Float32 model or an 80% Int8-20% Float32 model, the present invention, whether it is CP decomposition or Tucker decomposition, is superior to other methods in global model accuracy. These results show that the present invention can not only effectively handle the resource limitations of heterogeneous devices, but also significantly improve the accuracy of the model, especially in groups where low-bitwidth models account for the majority.
[0099] Table 1 Comparison of the present invention with the most advanced methods:
[0100]
[0101] See also Figure 5 and Figure 6 ,in, Figure 5 and Figure 6 The horizontal axis is the proportion of clients with float32 bit width in the clients participating in the training. In order to further verify the effect of the present invention under different bit width ratios, a more targeted experiment was designed and conducted. In the initial stage, all clients were deployed with Int8 bit width. Then the bit width of high-capacity clients was gradually deployed to Float32, and finally the global model accuracy under different Float32 bit width client ratios was observed. The results show that with the increase in the proportion of high-bit width clients, the global accuracy is significantly improved regardless of whether it is tensor filling based on CP decomposition or tensor filling based on tucker decomposition.
[0102] The above are only preferred embodiments of the present invention and are not intended to limit the present invention. For those skilled in the art, the present invention may have various modifications and variations. Any modification, equivalent replacement, improvement, etc. made within the spirit and principle of the present invention shall be included in the protection scope of the present invention.
Claims
1. A traversal federated learning method based on tensor filling, characterized in that: include: Deploy the bit width recovery component and the model segmentation component on the server side; Hierarchical tensor modeling of corresponding gradients based on device bit width; On the server side, the parameters of the same model layer in different gradients are classified and each layer is modeled as a tensor; Based on the tensor filling technology, the different bit width gradients corresponding to the tensors in the bit width recovery component are aligned, and all gradients are weighted aggregated, and a new global model is generated based on the aggregated weights; Based on the model segmentation component, a traversal window mechanism is constructed to segment the model. Through the traversal window mechanism, different layers of the global model are selected in the server and client for each round of training, so that the global model can be evenly trained in low-capacity devices until the global model converges; The hierarchical tensor modeling of the corresponding gradient based on the device bit width is specifically as follows: Identify model layers of the neural network that require bit-width-based tensor modeling; Based on the determined model layer, identify different data bit widths supported by the devices participating in the training, and determine the bit width dimension of the tensor accordingly; Map the model parameters trained by each device at this level to the corresponding position in the tensor according to its bit width; aggregate all device parameters under the same bit width to form a sub-tensor or element set; Parameter bit width mapping is performed independently for each layer of the neural network to ensure that the tensor can accurately reflect the parameter differences between devices with different bit widths on that layer; The model-based segmentation component is used to construct a traversal window mechanism segmentation model. Through the traversal window mechanism, different layers of the global model are selected in the server and the client for each round of training, so that the global model is evenly trained in low-capacity devices until the global model converges. Specifically, in each round of training of federated learning, the server uses the traversal window to extract sub-models of different capacities from the global model, and broadcasts them to clients with corresponding computing capabilities respectively; the client trains the received sub-models according to local data, and transmits the updated gradients of their heterogeneous sub-models to the server; the server summarizes the updated gradients and uses them to update the global model of the next round; the traversal window advances in each round, and loops through all parts of the global model in turn in different rounds; this process is iterated continuously until the global model reaches a state of uniform training after multiple trainings until convergence.
2. The traversal federated learning method based on tensor filling according to claim 1, characterized in that: The tensor-filling technology is based on aligning different bit-width gradients corresponding to tensors in the bit-width recovery component, specifically: the tensor size of the highest bit-width device is used as the standard tensor, and zero values are inserted into the tensor of the low bit-width device to reach the standard tensor; missing values in the low bit-width data are filled by tensor decomposition and tensor reconstruction, and the tensor of the low bit-width model parameters is aligned with the tensor shape of the full-precision model parameters.
3. The traversal federated learning method based on tensor filling according to claim 2, characterized in that: The tensor decomposition includes Tucker decomposition and CP decomposition; the Tucker decomposition is to convert a N The tensor is decomposed into a N The product of a core tensor of order and multiple factor matrices; for a N A tensor of order, of the form: (1); Among them, the elements in the tensor are represented as ,in is the element index, is the real number space, is the size of the tensor in each dimension. For example, for a three-dimensional tensor with a size of 10×20×30, N=3. , that is, there are 10 elements in the first dimension, 20 elements in the second dimension, and 30 elements in the third dimension. The tensor elements are represented as , and 1≤i1≤10, 1≤i2≤20, 1≤i3≤30; With a third-order tensor For example, its tucker decomposition is as follows: (2); in, represents the core tensor, is the dimension size of the core tensor, A, B, C represent factor matrices, represents the outer product; , and The factor matrices are A , B and C Column vector in ; is the core tensor G Elements in The CP decomposition is a method of decomposing the target tensor into the sum of multiple rank 1 tensors. N Rank Tensor , its CP is decomposed into the following form: (3) Among them, r is the rank of the decomposition selection, represents the outer product, Represents a factor matrix, each factor matrix belong , is the real number space, represents the size of the original tensor in the i-th dimension; i is a positive integer from 1 to N, and N is the order of the tensor; the tensor reconstruction is the inverse calculation of tensor decomposition, which is used to restore the tensor shape and approximate the original tensor; the CP decomposition attempts to find the minimum rank so that Approximate the original tensor as accurately as possible.
4. The traversal federated learning method based on tensor filling according to claim 3 is characterized in that: The tensor of the low-bit-width model parameters is aligned with the tensor shape of the full-precision model parameters, specifically: the low-bit-width model parameter elements are converted into multiple matrices, and expanded by adding column vectors and filling zero elements; the expanded matrices are decomposed and reconstructed by using CP decomposition or Tucker decomposition; After multiple rounds of iterations, a fully filled tensor is finally obtained; the newly added columns are merged into the original tensor and the shape is restored to complete the dequantization process of the model parameters.
5. The traversal federated learning method based on tensor filling according to claim 4, characterized in that: The method performs weighted aggregation on all gradients and generates a new global model based on the aggregated weights, specifically: assigning a weight to the gradient of each device or client based on the original bit width of the parameter; the weight is proportional to the accuracy of the bit width; and weighted aggregation of gradients from different devices or clients is performed using the determined weights.
6. The traversal federated learning method based on tensor filling according to claim 1, characterized in that: The server uses the traversal window to extract sub-models of different capacities from the global model and broadcasts them to clients with corresponding computing capabilities, specifically: Assuming a global model have L Layer, j The sub-models extracted by the server in round 1 include i Layer to The parameters of the layer; the traversal window moves forward one step in each round, skipping the model layer taken in the previous round; assuming that in the j In the round, the first n The client model is , then the traversal window starts from the global model Extract from i Layer to The method of the layer sub-model is expressed as: (4); in, Indicated in j In round, n Each client extracts the starting layer index of the sub-model; k is the size of the sub-model, that is, the number of layers of the sub-model. The traversal window will advance 1 layer per round during the training process and L The update position of the traversal window is expressed as: (5)。 7. The traversal federated learning method based on tensor filling according to claim 6, characterized in that: The client trains the received sub-models according to the local data, and transmits the updated gradients of their heterogeneous sub-models to the server. Specifically, since the local client only trains different sub-models in the global model in each round, the aggregation object is the heterogeneous sub-model extracted by the traversal window. For this heterogeneous model, the parameter part trained in this round is weighted and aggregated according to the aggregation strategy. The aggregation process is expressed as: (6); in, is the global model after aggregation, represents the learning rate, Indicated in j In round, n Client-to-submodel The updated gradient generated after training, N Represents the number of clients participating in the training, Representative n The local data of each client.
8. An ergodic federated learning system using the ergodic federated learning method according to any one of claims 1 to 7, characterized in that: include: A deployment module, wherein the deployment module deploys the bit width recovery component and the model segmentation component on the server side; A tensor modeling module, wherein the tensor modeling module performs hierarchical tensor modeling on corresponding gradients based on device bit width; On the server side, the parameters of the same model layer in different gradients are classified and each layer is modeled as a tensor; An alignment module, wherein the alignment module aligns different bit width gradients corresponding to tensors in the bit width recovery component based on tensor filling technology, performs weighted aggregation on all gradients, and generates a new global model based on the aggregated weights; The traversal module constructs a traversal window mechanism to segment the model based on the model segmentation component. Through the traversal window mechanism, different layers of the global model are selected in the server and the client for each round of training, so that the global model is evenly trained in the low-capacity device until the global model converges.
Citation Information
Patent Citations
Decentralization federated learning method under cross-domain heterogeneous scene
CN117556923A
Client data classification method, device and equipment based on longitudinal federated learning
CN117992840A