A method for federated learning of unbalanced data based on category distribution awareness

By establishing a correlation model between model parameters and loss function in federated learning, calculating the category distribution perception factor and constructing the cross-entropy loss function, the problem of global model accuracy degradation caused by unbalanced data distribution is solved, and the accuracy of object recognition is improved.

CN116662810BActive Publication Date: 2025-08-26HEFEI UNIV OF TECH
View PDF 2 Cites 0 Cited by

Patent Information

Application Number
CN202310678649.8
Authority / Receiving Office
CN · China
Patent Type
Patents(China)
Current Assignee / Owner
Filing Date
2023-06-07
Publication Date
2025-08-26
Estimated Expiration
2043-06-07

AI Technical Summary

Technical Problem

Under the unbalanced data distribution, the global model of federated learning has decreased accuracy and cannot accurately identify object categories.

Method used

By establishing a correlation model between model parameters and loss function during local model training, the category distribution perception factor is calculated, and the cross-entropy loss function perceived by category distribution is constructed to improve the accuracy of the global model.

Benefits of technology

Improve the global model accuracy of non-balanced data federated learning and improve the accuracy of object recognition.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN116662810B_ABST
    Figure CN116662810B_ABST
Patent Text Reader

Abstract

The present invention discloses a method for unbalanced data federated learning based on category distribution awareness, comprising: step 11, in which local users share a local model structure and calculate a category distribution awareness factor based on a local dataset; step 12, in which local users establish a correlation model between model parameters and a loss function in local model training, use the category distribution awareness factor and the absolute weight difference to construct a category distribution awareness cross-entropy loss function, update the local model, and upload the local model gradient to a server; step 13, in which the server aggregates the cross-entropy loss functions of local users to calculate the local model gradient and update the global model; step 14, in which local users download the global model from the server, and the server verifies whether the global model accuracy is greater than a preset threshold using an unbalanced dataset. If so, step 15 is executed; if not, steps 12 and 13 are repeated; and step 15 completes unbalanced data federated learning training. This method can improve the accuracy of the global model of unbalanced data federated learning.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the field of deep learning, and in particular to an unbalanced data federated learning method based on category distribution awareness. Background Art

[0002] With the abundance of computing and data resources, deep learning has rapidly developed and is widely applied in various real-world scenarios, such as object recognition, machine translation, and personalized recommendations. The successful application of deep learning requires a sufficiently large dataset. However, in real-world scenarios, the local data held by individual users is often insufficient. Therefore, it is necessary to fully utilize distributed local data. Federated learning requires distributed local users to collaboratively train a global model using their local data. The accuracy of the global model depends on the distribution of local user data. Due to the inconsistencies in local users' environments and data collection habits, local data often exhibits an imbalanced distribution. In the case of an imbalanced data distribution—that is, when the amount of data corresponding to each data category varies—the accuracy of the federated learning global model is often adversely affected. This also impacts the accuracy of the global model in object recognition.

[0003] In view of this, the present invention is proposed. Summary of the Invention

[0004] The purpose of the present invention is to provide an unbalanced data federated learning method based on category distribution awareness. By establishing an association model between model parameters and loss functions during local model training, calculating category distribution awareness factors, and constructing a category distribution awareness cross-entropy loss function, the present invention effectively solves the problem of decreased global model accuracy and inability to accurately identify object categories caused by the unbalanced distribution of local user data during the training process of the federated learning method.

[0005] The purpose of the present invention is achieved through the following technical solutions:

[0006] A method for federated learning of imbalanced data based on category distribution awareness, comprising:

[0007] Step 11: Local users share the network structure of the local model and calculate the category distribution perception factor based on the local dataset;

[0008] In step 12, the local user establishes an association model between model parameters and loss functions during the local model training process, uses the category distribution awareness factor obtained in step 11 and the absolute weight difference to construct a category distribution awareness cross entropy loss function, completes the local model update, and uploads the trained local model gradient as a shared parameter to the server;

[0009] Step 13: The server aggregates the local model gradients uploaded by local users and calculated using the category distribution-aware cross entropy loss function to complete the global model update.

[0010] Step 14: The local user communicates globally with the server and downloads the updated global model from the server. The server uses the unbalanced dataset to verify whether the accuracy of the global model is greater than a preset threshold. If so, step 15 is executed. If not, steps 12 and 13 are repeated.

[0011] Step 15: Complete the unbalanced data federated learning training process.

[0012] Compared with the existing technology, the unbalanced data federated learning method based on category distribution awareness provided by the present invention has the following beneficial effects:

[0013] By establishing an association model between model parameters and loss functions during local model training, calculating the category distribution-aware factor, and constructing a category distribution-aware cross-entropy loss function, the global model accuracy of unbalanced data federated learning is improved. BRIEF DESCRIPTION OF THE DRAWINGS

[0014] In order to more clearly illustrate the technical solutions of the embodiments of the present invention, the following briefly introduces the drawings required for use in the description of the embodiments. Obviously, the drawings described below are only some embodiments of the present invention. For ordinary technicians in this field, other drawings can be obtained based on these drawings without paying any creative work.

[0015] Figure 1 A flowchart of a method for federated learning of unbalanced data based on category distribution awareness provided in an embodiment of the present invention.

[0016] Figure 2 Schematic diagram showing the performance comparison of the unbalanced data federated learning method based on category distribution awareness provided in an embodiment of the present invention. DETAILED DESCRIPTION

[0017] The following is a clear and complete description of the technical solutions in the embodiments of the present invention in conjunction with the specific content of the present invention. Obviously, the embodiments described are only some embodiments of the present invention, not all embodiments, and do not constitute a limitation of the present invention. All other embodiments obtained by ordinary technicians in this field based on the embodiments of the present invention without making any creative efforts shall fall within the scope of protection of the present invention.

[0018] First, the following terms may be used in this article:

[0019] The term “and / or” means that either or both of them can be realized at the same time. For example, X and / or Y includes both “X” or “Y” and “X and Y”.

[0020] The terms "include," "comprises," "contains," "has," or other similar expressions should be interpreted as non-exclusive. For example, "including certain technical features (such as raw materials, components, ingredients, carriers, dosage forms, materials, dimensions, parts, components, mechanisms, devices, steps, procedures, methods, reaction conditions, processing conditions, parameters, algorithms, signals, data, products, or manufactured articles)" should be interpreted as including not only the technical features explicitly listed, but also other technical features known in the art that are not explicitly listed.

[0021] The term "consisting of" excludes any technical features not explicitly listed. If used in a claim, this term renders the claim closed, excluding any technical features other than those explicitly listed, except for conventional impurities associated with them. If this term appears only in a clause of a claim, it limits only the elements explicitly listed in that clause; elements listed in other clauses are not excluded from the claim as a whole.

[0022] Unless otherwise specified or limited, the terms "mounted," "connected," "connect," and "fixed" should be interpreted broadly. For example, they can refer to fixed, detachable, or integral connections; mechanical or electrical connections; direct or indirect connections through an intermediary; and internal communication between two components. Those skilled in the art will understand the specific meanings of the above terms in this document based on specific circumstances.

[0023] The terms "center", "longitudinal", "lateral", "length", "width", "thickness", "up", "down", "front", "back", "left", "right", "vertical", "horizontal", "top", "bottom", "inside", "outside", "clockwise", "counterclockwise", etc., indicating the orientation or position relationship, are based on the orientation or position relationship shown in the accompanying drawings and are only for the convenience and simplification of description, and do not explicitly or implicitly indicate that the device or element referred to must have a specific orientation, be constructed and operate in a specific orientation, and therefore should not be understood as a limitation to this document.

[0024] The following describes in detail the unbalanced data federated learning method based on class distribution awareness provided by the present invention. Any content not described in detail in the examples of the present invention belongs to the prior art known to professionals in the field. Where specific conditions are not specified in the examples of the present invention, the methods were performed according to conventional conditions in the field or the conditions recommended by the manufacturer. Reagents or instruments used in the examples of the present invention, where the manufacturer is not specified, are all commercially available conventional products.

[0025] like Figure 1 As shown, an embodiment of the present invention provides an unbalanced data federated learning method based on category distribution awareness, comprising the following steps:

[0026] Step 11: Local users share the network structure of the local model and calculate the category distribution perception factor based on the local dataset;

[0027] In step 12, the local user establishes an association model between model parameters and loss functions during the local model training process, uses the category distribution awareness factor obtained in step 11 and the absolute weight difference to construct a category distribution awareness cross entropy loss function, completes the local model update, and uploads the trained local model gradient as a shared parameter to the server;

[0028] Step 13: The server aggregates the local model gradients uploaded by local users and calculated using the category distribution-aware cross entropy loss function to complete the global model update.

[0029] Step 14: The local user communicates globally with the server and downloads the updated global model from the server. The server uses the unbalanced dataset to verify whether the accuracy of the global model is greater than a preset threshold. If so, step 15 is executed. If not, steps 12 and 13 are repeated.

[0030] Step 15: Complete the unbalanced data federated learning training process.

[0031] Preferably, in step 11 of the above method, if there are local users , then each local user Possession includes local datasets of data , each local dataset contain Due to the inconsistency between local users' environments and data collection habits, local data usually presents an unbalanced distribution, that is, the number of data corresponding to each data category varies;

[0032] The network structure of the local model shared by local users is ResNet-20. The category distribution perception factor is calculated based on the local dataset using the following formula: for:

[0033] (1);

[0034] in, For the Class data in local dataset the proportion of is a parameter with a value of 1.02; The value range is .

[0035] Preferably, in step 12 of the above method, the local user establishes an association model between model parameters and loss functions during local model training in the following manner, uses the category distribution awareness factor and the absolute weight difference to construct a category distribution awareness cross entropy loss function, completes the local model update, and uses the trained local model gradient as a shared parameter for uploading to the server, including:

[0036] Step 121, in each local model update process, the parameters of the local model of each local user are calculated by Updated to , the category distribution-aware cross entropy loss function used by the local model Expressed as Taylor series:

[0037] (2);

[0038] in, The data in the local dataset;

[0039] Ignoring the high-order terms in formula (2), the association model between the model parameters and the category distribution-aware cross entropy loss function during the local model training process established by the local user is:

[0040] (3);

[0041] Step 122: Each local user calculates the first The absolute weight difference and for:

[0042] (4);

[0043] in, For the The parameters of the last fully connected layer in the local model during the local model update; For the The parameters of the last fully connected layer in the local model when the local model is updated are parameters; According to formula (3), as the local model converges, the absolute weight difference and Approaching zero;

[0044] In step 123, the local user uses the category distribution perception factor according to the following formula (5): and the absolute weight difference The cross entropy loss function for constructing category distribution awareness is:

[0045] (5);

[0046] in, For the Local model parameters during the local model update; For data Corresponding to The true category label of the class; For data Corresponding to The predicted class label of the class; is the category-aware weight, Expressed as:

[0047] (6);

[0048] In formula (6), is the normalized category distribution perception factor, ;

[0049] Step 124, The cross entropy loss function of local users using category distribution awareness Complete the The local model update is expressed as:

[0050] (7);

[0051] in, For the Local users completed the Local model parameters during the local model update; For the Local users completed the Local model parameters during the local model update; For the Local users completed the The local model gradient is calculated by the category distribution-aware cross entropy loss function during the local model update; is the updated learning rate of the local model;

[0052] Step 125: Each local user repeats steps 122), 123), and 124 until the number of local model updates reaches the number of local model update rounds. Until the time, the local model gradient of the training is completed Uploaded to the server as shared parameters.

[0053] Preferably, in step 13 of the above method, the server aggregates the local model gradients uploaded by each local user and calculated by the category distribution-aware cross entropy loss function to complete the global model update, which is expressed as:

[0054] (8);

[0055] in, For the Global model parameters during sub-global communication; For the Global model parameters during sub-global communication; Update the learning rate for the global model; The number of local data for all local users; For the The first global communication The first data uploaded by local users is calculated by the category distribution-aware cross entropy loss function. The local model gradient of each local user.

[0056] Preferably, in step 14 of the above method, During the global communication, all local users download the updated global model from the server The server uses an unbalanced dataset that is independent and identically distributed with the local dataset to verify whether the accuracy of the global model is greater than a preset threshold (the preset threshold is preferably 0.5). If so, step 15 is executed. Otherwise, steps 12 and 13 are repeated.

[0057] In summary, the method of the embodiment of the present invention improves the global model accuracy of unbalanced data federated learning by establishing an association model between model parameters and loss functions during local model training, calculating category distribution-aware factors, and constructing a category distribution-aware cross-entropy loss function.

[0058] In order to more clearly demonstrate the technical solution and technical effects provided by the present invention, the unbalanced data federated learning method based on category distribution awareness provided by an embodiment of the present invention is described in detail below with reference to a specific embodiment.

[0059] Example 1

[0060] The embodiment of the present invention provides an unbalanced data federated learning method based on category distribution awareness for object recognition, which specifically includes:

[0061] Step 1: The local user obtains image data of the object to be identified;

[0062] In step 2, the local user downloads a global object recognition model pre-trained by unbalanced data federated learning from the server, recognizes the image data of the object to be recognized obtained in step 1, and obtains the category of the object.

[0063] like Figure 1 As shown in FIG, the unbalanced data federated learning training process for training the global object recognition model in step 2 above mainly includes the following steps:

[0064] In step 11, local users share the network structure of the local model and calculate the category distribution perception factor based on the local dataset.

[0065] The preferred implementation of this step 11 is as follows:

[0066] Step 111) Federated Learning local users , each local user Possession includes local datasets of data , each local dataset contain Due to the inconsistency between local users' environments and data collection habits, local data usually presents an unbalanced distribution, that is, the number of data corresponding to each data category varies;

[0067] Step 112) The local user shares a local model with a ResNet-20 network structure and calculates the category distribution perception factor based on the local data set. , expressed as:

[0068] (1);

[0069] in, For the Class data in local dataset The proportion of is a parameter with a value of 1.02, The value range is .

[0070] In step 12, the local user establishes an association model between model parameters and loss functions during local model training, uses the category distribution awareness factor and the absolute weight difference to construct a category distribution awareness cross entropy loss function, completes the local model update, and uploads the trained local model gradient as a shared parameter to the server.

[0071] The preferred implementation of this step 12 is as follows:

[0072] Step 121) For each local user, the parameters of the local model are calculated by Updated to By using Taylor series, the local model uses a category distribution-aware cross entropy loss function Expressed as:

[0073] (2);

[0074] in, is the data in the local dataset. Ignoring the high-order terms in formula (2), the local user establishes an association model between the model parameters and the category distribution-aware cross entropy loss function during the local model training process, which is expressed as:

[0075] (3);

[0076] Step 122) Each local user calculates the The absolute weight difference and , expressed as:

[0077] (4);

[0078] in, For the The parameters of the last fully connected layer in the local model when the local model is updated. For the The parameters of the last fully connected layer in the local model when the local model is updated are parameters.

[0079] Step 123) Local users use category distribution perception factors and the absolute weight difference Construct a category distribution-aware cross entropy loss function, expressed as:

[0080] (5);

[0081] in, For the The local model parameters when the local model is updated, For data Corresponding to The true category label of the class, For data Corresponding to The predicted class label for the class.

[0082] In formula (5), the category-aware weight Expressed as:

[0083] (6);

[0084] In formula (6), the normalized category distribution perception factor Expressed as .

[0085] Step 124) For local users, using the category distribution-aware cross entropy loss function Complete the The local model update is expressed as:

[0086] (7);

[0087] in, For the Local users completed the The local model parameters when the local model is updated, For the Local users completed the The local model parameters when the local model is updated, For the Local users completed the The local model gradient is calculated by the category distribution-aware cross entropy loss function during the local model update. Update the learning rate for the local model.

[0088] Step 125) Each local user repeats steps (2), (3) and (4) until the number of local model updates reaches the number of local model update rounds. Until the time, the local model gradient of the training is completed As a shared parameter for uploading to the server.

[0089] In step 13, the server aggregates the local model gradients uploaded by local users and calculated by the category distribution-aware cross entropy loss function to complete the global model update.

[0090] Among them, the global model update is expressed as:

[0091] (8);

[0092] in, For the Global model parameters during sub-global communication, For the Global model parameters during sub-global communication, Update the learning rate for the global model, is the number of local data of all local users, For the The first global communication The gradient is calculated by the category distribution-aware cross entropy loss function uploaded by local users.

[0093] In step 14, the local user communicates globally with the server and downloads the updated global model from the server. The server uses the unbalanced dataset to verify whether the accuracy of the global model is greater than a preset threshold. If so, step 15 is executed. Otherwise, steps 12 and 13 are repeated.

[0094] The preferred implementation of this step 14 is as follows:

[0095] In the During the global communication, all local users download the updated global model from the server , the server uses an unbalanced dataset that is independent and identically distributed with the local dataset to verify whether the accuracy of the global model is greater than a preset threshold. If so, step 15 is executed. If not, steps 12 and 13 are repeated.

[0096] Step 15: Complete the unbalanced data federated learning training process.

[0097] The unbalanced data federated learning method based on category distribution awareness in an embodiment of the present invention improves the global model accuracy of unbalanced data federated learning by establishing an association model between model parameters and loss functions during local model training, calculating category distribution awareness factors, and constructing a category distribution awareness cross-entropy loss function.

[0098] As you can see, for object recognition, the datasets used in the above training are all object recognition datasets, and the data in the datasets are all image data used for object recognition. The local models of local users are all local object recognition models, and the global models on the server are all global object recognition models. As the accuracy of the global model is improved, the accuracy of the final object recognition is also correspondingly improved.

[0099] To test the global model accuracy of the class distribution-aware unbalanced data federated learning method of this embodiment, this federated learning method is compared with an impact-balanced unbalanced data federated learning method. The impact-balanced unbalanced data federated learning method is denoted as Method A.

[0100] The global model accuracy was tested on the object recognition datasets MNIST, Fashion-MNIST, and CIFAR-10, using stochastic gradient descent as the basic model optimization algorithm. The number of local users was set to 5, the number of global communication rounds was set to 30, the number of local model update rounds was set to 100, the random sampling dataset size was set to 64, and the local and global model update learning rates were set to 0.01. The data imbalance ratio ρ was defined as the ratio of the number of data points corresponding to the minority class to the number of data points corresponding to the majority class in the dataset. Ten tests were conducted on the global model accuracy under the given experimental parameter settings. The test results are shown in Table 1.

[0101] Table 1 Global model accuracy of different imbalanced data federated learning methods

[0102] .

[0103] Table 1 shows the global model accuracy of different imbalanced data federated learning methods, where each global model accuracy is expressed as the mean and standard deviation. The class distribution-aware imbalanced data federated learning method of the present invention can achieve higher global model accuracy in most cases.

[0104] Figure 2 Shows the number of categories Class-aware weight The horizontal axis is the sum of absolute weight differences, and the vertical axis is the class-wise weight. Figure 2 It can be seen that the method of the present invention can set a larger weight for the data corresponding to the minority class, thereby alleviating the overfitting problem.

[0105] Those skilled in the art will appreciate that all or part of the processes in the above-described method embodiments can be implemented by instructing related hardware through a program. The program can be stored in a computer-readable storage medium. When executed, the program can include the processes in the above-described method embodiments. The storage medium can be a magnetic disk, an optical disk, a read-only memory (ROM), or a random access memory (RAM).

[0106] The above description is only a preferred embodiment of the present invention, but the scope of protection of the present invention is not limited thereto. Any changes or substitutions that can be easily thought of by any person skilled in the art within the technical scope disclosed in the present invention should be included in the scope of protection of the present invention. Therefore, the scope of protection of the present invention should be based on the scope of protection of the claims. The information disclosed in the background technology section of this article is only intended to deepen the understanding of the overall background technology of the present invention, and should not be regarded as an admission or any form of implication that the information constitutes prior art already known to those skilled in the art.

Claims

1. A method for unbalanced data federated learning based on category distribution awareness, characterized in that: include: Step 11: Local users share the network structure of the local model and calculate the category distribution perception factor based on the local dataset; In step 12, the local user establishes an association model between model parameters and loss functions during the local model training process, uses the category distribution awareness factor and absolute weight difference obtained in step 11 to construct a category distribution awareness cross entropy loss function, completes the local model update, and uploads the trained local model gradient as a shared parameter to the server; specifically: Step 121, in each local model update process, the parameters of the local model of each local user are calculated by Updated to , the category distribution-aware cross entropy loss function used by the local model Expressed as a Taylor series; Step 122: Each local user calculates the first The absolute weight difference and for: (4); in, For the The parameters of the last fully connected layer in the local model during the local model update; For the The parameters of the last fully connected layer in the local model when the local model is updated are parameters; as the local model converges, the absolute weight difference and Approaching zero; In step 123, the local user uses the category distribution perception factor according to the following formula (5): and the absolute weight difference The cross entropy loss function for constructing category distribution awareness is: (5); in, For the Local model parameters during the local model update; For data Corresponding to The true category label of the class; For data Corresponding to The predicted class label of the class; is the category-aware weight; Step 124, The cross entropy loss function of local users using category distribution awareness Complete the The local model update is expressed as: (7); in, For the Local users completed the Local model parameters during the local model update; For the Local users completed the Local model parameters during the local model update; For the Local users completed the The local model gradient is calculated by the category distribution-aware cross entropy loss function during the local model update; is the updated learning rate of the local model; Step 125: Each local user repeats steps 122), 123), and 124 until the number of local model updates reaches the number of local model update rounds. Until the time, the local model gradient of the training is completed Upload to the server as shared parameters; Step 13: The server aggregates the local model gradients uploaded by local users and calculated using the category distribution-aware cross entropy loss function to complete the global model update. Step 14: The local user communicates globally with the server and downloads the updated global model from the server. The server uses the unbalanced dataset to verify whether the accuracy of the global model is greater than a preset threshold. If so, step 15 is executed. If not, steps 12 and 13 are repeated. Step 15: Complete the unbalanced data federated learning training process.

2. The unbalanced data federated learning method based on category distribution awareness according to claim 1 is characterized in that: In step 11, if there are local users , then each local user Possession includes local datasets of data , each local dataset contain The data of the local dataset is unbalanced; The network structure of the local model shared by local users is ResNet-20. The category distribution perception factor is calculated based on the local dataset using the following formula: for: (1); in, For the Class data in local dataset the proportion of is a parameter with a value of 1.02; The value range is .

3. The unbalanced data federated learning method based on category distribution awareness according to claim 1 or 2, characterized in that: In step 121, the category distribution-aware cross entropy loss function used by the local model is Expressed as Taylor series: (2); in, The data in the local dataset; Ignoring the high-order terms in formula (2), the association model between the model parameters and the category distribution-aware cross entropy loss function during the local model training process established by the local user is: (3); In step 123, the category-aware weight Expressed as: (6); In formula (6), is the normalized category distribution perception factor, .

4. The unbalanced data federated learning method based on category distribution awareness according to claim 1 or 2, characterized in that: In step 13, the server aggregates the local model gradients uploaded by each local user and calculated by the category distribution-aware cross entropy loss function to complete the global model update, which is expressed as: (8); in, For the Global model parameters during sub-global communication; For the Global model parameters during sub-global communication; Update the learning rate for the global model; The number of local data for all local users; For the The first global communication The first data uploaded by local users is calculated by the category distribution-aware cross entropy loss function. The local model gradient of each local user.

5. The unbalanced data federated learning method based on category distribution awareness according to claim 1 is characterized in that: In the step 14, During the global communication, all local users download the updated global model from the server , the server uses an unbalanced dataset that is independent and identically distributed with the local dataset to verify whether the accuracy of the global model is greater than a preset threshold. If so, step 15 is executed. If not, steps 12 and 13 are repeated.

6. The unbalanced data federated learning method based on category distribution awareness according to claim 1 or 5, characterized in that: The preset threshold value is 0.5.

Citation Information

Patent Citations

  • Industrial Internet of Things privacy protection system and method based on federal learning

    CN114417417A

  • Federal learning method for long-tail heterogeneous data

    CN114429219A