Federal learning client selection method based on reinforcement learning and federal learning system

Through reinforcement learning and mean drift clustering optimization client selection, the problems of data heterogeneity and resource consumption imbalance in federated learning are solved, and the convergence speed and generalization capabilities of the global model are improved, which is suitable for edge computing environments.

CN120297439APending Publication Date: 2025-07-11HOHAI UNIV
View PDF 0 Cites 6 Cited by

Patent Information

Application Number
CN202510392233.9
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-03-31
Publication Date
2025-07-11

AI Technical Summary

Technical Problem

The existing federated learning method fails to effectively solve the problems of client data heterogeneity, resource constraints and communication overhead, resulting in slow convergence speed and degradation of performance in the global model.

Method used

The client selection method based on reinforcement learning is adopted, and clients are grouped through mean drift clustering, and the multi-agent reinforcement learning framework is used to dynamically select clients to participate in training. Combining multi-dimensional state information and dynamic exploration strategies, client selection is optimized to balance communication and energy consumption.

Benefits of technology

It improves the training efficiency and performance of the global model, reduces computing and communication overhead, enhances the generalization ability of the model under different data distributions, and is suitable for resource-constrained edge computing environments.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120297439A_ABST
    Figure CN120297439A_ABST
Patent Text Reader

Abstract

The invention discloses a federated learning client selection method based on reinforcement learning and a federated learning system. The method comprises the following steps: a client performs local detection training and collects loss, delay data and other related data; the central server divides the clients by using mean shift clustering based on data distribution of the clients, and regards each cluster as an independent agent; the intelligent agent dynamically selects a client based on the multi-dimensional state information, and optimization selection is carried out by adopting an exploration strategy; the client uploads model update after local training; and the central server carries out aggregation updating, and the intelligent agent optimizes client selection through a multi-target reward function according to a feedback adjustment strategy, maximizes a model convergence speed and balances communication and calculation overhead. By using the method and the system of the invention, under challenged environments of data heterogeneity, computing resource limitation, communication delay and the like, the client can be intelligently selected, the training efficiency of federated learning is optimized, and the performance and generalization ability of a global model are improved.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the technical field of distributed computing, and particularly relates to a method for selecting a federated learning client based on reinforcement learning and a federated learning system. Background Art

[0002] With the continuous development of artificial intelligence and big data technologies, Federated Learning (FL), as a new type of distributed machine learning method, has been attracting increasing attention. Under the federated learning framework, multiple distributed clients (such as smartphones, Internet of Things devices, etc.) cooperate to train a shared global model without directly exchanging local data. The client only needs to send the model parameters obtained from local training to the server, and the server aggregates these model parameters to update the global model. Although federated learning has significant advantages in data privacy protection, it also faces many technical challenges in practical applications, especially problems such as heterogeneity among clients, resource constraints, and high communication overhead.

[0003] In a federated learning system, the data of each client usually has significant heterogeneity (Non-Independent and Identically Distributed, Non-IID). In practical applications, client data is often generated by different users, and the data distribution varies greatly. For example, social media data on smartphones is very different from physiological data collected by health monitoring devices in terms of feature distribution, data volume, data quality, etc. This heterogeneity leads to differences in the convergence speed and performance of the global model on different clients, thereby affecting the training effect of the global model.

[0004] In addition to the data heterogeneity problem, federated learning also faces problems of frequent communication and resource constraints. Since after each training, the client needs to upload the updated model parameters to the central server, frequent communication will bring significant bandwidth overhead. In addition, each client in federated learning usually has limited computing power and energy consumption budget. Especially for mobile devices, the computing and communication consumption during the training process may cause the device performance to decline or even the battery to run out quickly. Therefore, how to balance communication overhead, computing power, and energy consumption has become a key issue in improving the efficiency of federated learning.

[0005] In response to the above problems, the optimization of the client selection method has become an important research topic in federated learning. Traditional client selection methods usually rely on preset rules, such as randomly selecting some clients or choosing clients with high computing power to participate in training. However, these traditional methods fail to fully consider the heterogeneity of client data distribution, different resource consumption, and the balance of communication overhead, and often cannot achieve the optimal effect. Furthermore, traditional methods often ignore the long-term cumulative reward problem in federated learning. Local optimization cannot effectively guide global optimization, which in turn leads to a slow convergence rate of the global model and even a decline in performance.

[0006] Reinforcement Learning (RL), as an adaptive decision-making optimization method, has been widely applied in many fields in recent years. In federated learning, reinforcement learning can dynamically adjust the client selection strategy to adapt to different data distributions and performance requirements. Different from traditional methods, reinforcement learning learns through real-time interaction and adjusts the client selection strategy according to environmental feedback (such as communication overhead, energy consumption, training effect, etc.). Specifically, reinforcement learning can optimize the selection of clients in real time by evaluating the contribution of each client, effectively alleviating the impact of data heterogeneity on model convergence, balancing communication and energy consumption, and ultimately improving the performance of the global model. Summary of the Invention

[0007] Object of the Invention: The objective of the present invention is to optimize the client selection strategy in federated learning through a reinforcement learning algorithm. A method based on reinforcement learning is proposed for the problems of client data heterogeneity, communication overhead, and energy consumption. The clients are grouped through data distribution clustering, and then appropriate clients are intelligently selected in each cluster to participate in training, so as to reduce the impact of data heterogeneity on model convergence and balance communication and energy consumption. The ultimate goal is to improve the performance and training efficiency of the global model, while ensuring its efficiency and generalization ability in resource-constrained environments.

[0008] To achieve the above object of the invention, the technical solution of the present invention is as follows:

[0009] In a first aspect, a method for selecting clients in federated learning based on reinforcement learning includes the following steps:

[0010] The client conducts local probing training and state collection, including: each client conducts local probing training based on the current global model, calculates the local model update, and collects the probing training loss, the corrected loss corrected by the control variable, and the probing training delay; records the client's historical communication delay, training delay, communication energy consumption, local dataset size, participation frequency, recent participation interval, and training round index;

[0011] The central server clusters and partitions the clients, including: Based on the local data distribution of the clients, using the mean shift clustering algorithm to partition the clients into multiple data - homogeneous clusters, generating cluster identifiers and cluster statistical features; Each cluster corresponds to a reinforcement learning agent, which is responsible for dynamically selecting clients within the cluster to participate in training;

[0012] The central server performs agent state modeling and action decision - making, including: The agent dynamically selects clients based on multi - dimensional state information including detected training loss, corrected loss, control variable difference, historical communication delay, training delay, communication energy consumption, local dataset size, participation frequency, recent participation interval, training round index, cluster identifier, and cluster statistical features; Adopting a dynamic exploration strategy and frequency - limit mechanism to optimize client selection;

[0013] The clients perform local training and model uploading, including: The selected clients perform local training, calculate model updates, and upload them to the central server;

[0014] The central server performs global model aggregation and feedback, including: The central server aggregates the model updates uploaded by the clients, optimizes the global model, and broadcasts it to each cluster; The agent adjusts the strategy according to the feedback of the global model, and optimizes the client selection strategy through a multi - objective reward function.

[0015] Furthermore, the mean shift clustering algorithm calculates the mean within the feature space of each client, iteratively updates the clustering centers of each client, and finally converges to a set of representative clusters;

[0016] The cluster statistical features include:

[0017] Data - class distribution entropy: Measuring the distribution diversity of data classes within the cluster;

[0018] Data volume variance: Measuring the unevenness of the data volumes of clients within the cluster;

[0019] Client resource mean: Calculating the average of the computing capabilities and communication energy consumptions of clients within the cluster.

[0020] Furthermore, the state of the agent is as follows:

[0021] s t =[l t ,l′ t ,De t ,T comm,t ,T train,t ,E t ,D t ,G t ,F t ,C t ,S stats,t ​

[0022] Wherein:

[0023] l t is the detection training loss of client t;

[0024] l′ t is the loss of client t after control variable correction;

[0025] De t is the difference between the local control variable of client t and the global model control variable;

[0026] T comm,t is the communication delay of client t;

[0027] T train,t is the local training time of client t;

[0028] E t is the communication energy consumption of client t;

[0029] D t is the size of the local dataset of client t;

[0030] G t is the frequency of client t participating in training;

[0031] F t is the time interval since client t last participated in training;

[0032] C t is the identifier of the cluster to which client t belongs;

[0033] S stats,t is the statistical feature of the cluster where client t is located.

[0034] Furthermore, the difference De between the local control variable of client t and the global model control variable t is calculated as follows:

[0035] De t =||c t -c||2

[0036] Wherein, c t is the local control variable of client t, and c is the global control variable.

[0037] Furthermore, assuming that each cluster C k contains M k clients, the action space A of the agent k is expressed as:

[0038]

[0039] The action of the agent is represented as a binary vector, where each element a t ∈ {0, 1}, indicating whether to select the t-th client in cluster k.

[0040] Furthermore, the action selection strategy includes:

[0041] Dynamic exploration strategy: ε-greedy decay in phases, where ε is the decay factor;

[0042] Forced exploration mechanism: For clients that have not participated in training for more than N rounds, select them with a specified probability;

[0043] Frequency limit: If the number of participations of a client in the past specified rounds ≥ the specified threshold, force it to skip one round.

[0044] Furthermore, each agent uses a multi-agent reinforcement learning algorithm to select clients. The algorithm updates the policy based on the multi-objective reward function feedback from the global model. The reward function is as follows:

[0045]

[0046] Where:

[0047] ΔAcc is the improvement in the global model test accuracy;

[0048] H t is the maximum communication and training delay in the current round;

[0049] De k is the difference in client control variables;

[0050] Num(K t ) is the number of clusters covered in the current round;

[0051] ω1, ω2, ω3, β are weight coefficients.

[0052] Furthermore, the optimization objective is represented by the following formula:

[0053]

[0054] Where:

[0055] A is the client selection policy matrix;

[0056] Acc(T) is the test accuracy of the global model in round T;

[0057] H t is the total processing delay in round t;

[0058] B t is the total communication energy consumption in round t.

[0059] In a second aspect, a federated learning system includes a number of clients and a central server. The clients are configured to:

[0060] Perform local probing training based on the current global model, calculate local model updates, and collect probing training losses, corrected losses corrected by control variables, and probing training delays; record the client's historical communication delays, training delays, communication energy consumption, local dataset sizes, participation frequencies, recent participation intervals, and training round indices; upload the status data to the central server;

[0061] And when selected by the central server, perform local training, calculate model updates, and upload them to the central server;

[0062] The central server is configured to:

[0063] Based on the local data distribution of the clients, use the mean shift clustering algorithm to divide the clients into multiple data homogeneous clusters, generate cluster identifiers and cluster statistical features; each cluster corresponds to a reinforcement learning agent responsible for dynamically selecting clients within the cluster to participate in training;

[0064] Perform agent state modeling and action decision-making. The agent state space includes probing training losses, corrected losses, control variable differences, historical communication delays, training delays, communication energy consumption, local dataset sizes, participation frequencies, recent participation intervals, training round indices, cluster identifiers, and cluster statistical features. The agent dynamically selects clients based on multi-dimensional state information, and adopts a dynamic exploration strategy and a frequency limit mechanism to optimize client selection;

[0065] And aggregate the model updates uploaded by the clients, optimize the global model, and broadcast it to each cluster; and adjust the strategy according to the feedback of the global model, and optimize the client selection strategy through a multi-objective reward function.

[0066] Beneficial effects: (1) The present invention proposes a method for selecting federated learning clients based on reinforcement learning. By using the mean shift clustering algorithm and the multi-agent reinforcement learning (MARL) framework, clients participating in training are intelligently selected. This method can effectively solve problems such as data heterogeneity, resource limitations, and communication delays among clients in federated learning, thereby improving the training efficiency and performance of the global model. Through multi-dimensional state modeling and dynamic exploration strategies, the selection of clients is optimized, avoiding the limitations of traditional methods and reducing unnecessary computational and communication overheads. (2) The present invention dynamically adjusts the client selection strategy through reinforcement learning, solving the problem that traditional client selection methods do not adequately consider the balance between data heterogeneity and resource consumption. Through detailed state information (such as training loss, communication energy consumption, control variable differences, etc.), the system can accurately select the optimal client for training, effectively improving the convergence speed of the global model and enhancing the generalization ability of the model under different data distributions. (3) The multi-objective reward mechanism of the present invention considers multiple factors such as model accuracy improvement, communication energy consumption, and training delay. Without increasing the additional communication burden, it balances the overheads of computational and communication resources, significantly improving the efficiency of federated learning. This method is particularly suitable for edge computing environments and can perform efficient model training when the computing power of devices is limited. It has strong adaptability and versatility, providing a new solution for the federated learning application of intelligent devices. Brief Description of the Drawings

[0067] Figure 1 is a flowchart of the method for selecting federated learning clients based on reinforcement learning according to the present invention;

[0068] Figure 2 is a schematic diagram of client clustering and cluster division;

[0069] Figure 3 is a decision flowchart of the reinforcement learning agent. Detailed Embodiments

[0070] The technical solutions of the present invention will be further described below with reference to the accompanying drawings.

[0071] The present invention proposes a method for selecting federated learning clients based on reinforcement learning. The purpose of this method is to optimize the federated learning process by intelligently selecting clients participating in training and solve problems such as data heterogeneity, limited computing resources, and communication delays. The implementation process is divided into main steps such as client detection training and state collection, client clustering and cluster division, and agent state modeling and action decision-making, ensuring that the state information, resource consumption of clients, and optimization of the global model are fully considered in each step.

[0072] The present invention uses a reinforcement learning algorithm to dynamically adjust the client selection strategy, and combines the mean shift clustering algorithm to divide the clients into multiple clusters according to data distribution, thereby optimizing the client selection, improving the convergence speed of the global model, and reducing the communication overhead and computing burden. Figure 1 , the method of the present invention comprises the following steps:

[0073] Step (1): client detection training and status collection.

[0074] In the present invention, the client needs to perform local detection training in each round of federated learning to calculate the update of the local model and collect various state information related to the training process. This state information will be used for subsequent decision-making of the intelligent agent to help it select the most appropriate client to participate in the training.

[0075] Probe training: Each client performs local training based on the current global model, i.e., "probe training", in order to generate local model updates. The process of probe training is similar to traditional local training, except that during the local training process, the client will use the global model as the initialization, perform a certain number of rounds of training, and generate local model updates. During this process, the client calculates and records the probe training loss (indicating the degree of match between the local data and the current global model), as well as the corrected loss (the loss corrected by the control variables, taking into account the potential deviations in the local training process).

[0076] Collecting status information: In addition to the training loss, the client also collects a series of other training-related status information, including:

[0077] Training latency: The time required for the client to perform local training.

[0078] Communication energy consumption: The communication energy consumed by the client to upload local model updates.

[0079] Local dataset size: The normalized proportion of the client's local data volume in the cluster to which it belongs, used to evaluate the relative importance of the client's data volume.

[0080] Participation frequency: The frequency at which the client participated in training in the past several rounds (using exponential decay mean).

[0081] Recent participation interval: The difference between the last round in which the client participated in training and the current round, used to evaluate the activeness of the client's participation.

[0082] Historical communication delay: The delay of the client uploading and downloading model parameters in the past several rounds (Δ≥5 rounds).

[0083] Control variable difference: client t local control variable c tGradient correction term difference De between the global control variable c t , quantize the gradient deviation.

[0084] De t = ||c t - c||2

[0085] Summary of status information: Each client conducts exploratory training locally based on the current global model, calculates model updates, and collects the above information. This status information is used to describe the current situation of the client for subsequent agent decision-making.

[0086] Step (2), client clustering and cluster division.

[0087] After local training, the client uploads the model update to the server. The server uses the mean shift clustering algorithm to analyze the model updates of the clients, divides the clients with similar data distributions into homogeneous clusters, assigns a unique identifier to each cluster, and calculates the cluster statistical features.

[0088] The local data distributions of clients usually have significant heterogeneity, which can affect the convergence and performance of the global model. Therefore, the present invention uses the mean shift clustering algorithm to cluster clients according to data distributions and regards each cluster as an independent agent. The clients within each cluster have similar characteristics and data distributions, and the training efficiency can be improved through the optimization strategy of the agent within the cluster. A schematic diagram of client clustering and cluster division is as Figure 2 .

[0089] Mean shift clustering algorithm: The present invention uses the mean shift clustering algorithm to automatically group clients. The goal of clustering is to group clients with similar data distributions into the same group. Specifically, the mean shift clustering algorithm calculates the mean within the feature space of each client, iteratively updates the clustering center of each client, and finally converges to a set of representative clusters. The number of clusters does not need to be set in advance during the clustering process, so it can be adaptively divided according to the complexity of the data distribution.

[0090] Statistical features of the cluster: Each clustering cluster generates a cluster identifier to distinguish different clusters. And according to the data characteristics of the clients within the cluster, the statistical features of the cluster are calculated, and these features are used to further optimize the client selection strategy. The cluster statistical features include:

[0091] Entropy of data category distribution: Measures the distribution diversity of data categories within the cluster.

[0092] Variance of data volume: Measures the unevenness of the data volumes of clients within the cluster.

[0093] Mean of client resources: Calculates the average of the computing capabilities and communication energy consumption of clients within the cluster.

[0094] Step (3), agent state modeling and action decision-making.

[0095] In the present invention, each clustering cluster corresponds to an agent. The task of the agent is to dynamically select in-cluster clients to participate in training according to the multi-dimensional state information of each client. The central server, as the global coordinator, is responsible for collecting the model updates and state information of the clients and performing clustering and agent decision-making. Through reinforcement learning, the agent continuously adjusts its strategy to optimize client selection and improve the performance of the global model. The state of the agent is as follows:

[0096] s t =[l t ,l′ t ,De t ,T comm,t ,T train,t ,E t ,D t ,G t ,F t ,C t ,S stats,t

[0097] Where:

[0098] l t : The probing training loss of client t.

[0099] l′ t : The loss of client t after being corrected by the control variable.

[0100] De t : The difference between the local control variable of client t and the global model control variable.

[0101] T comm,t : The communication delay of client t.

[0102] T train,t : The local training time of client t.

[0103] E t : The communication energy consumption of client t.

[0104] D t : The size of the local dataset of client t.

[0105] G t : The frequency of client t participating in training.

[0106] F t : The time interval since client t last participated in training.

[0107] C t : The identifier of the cluster to which client t belongs. ​

[0108] S stats,t : Statistical features of the cluster where client t is located.

[0109] Action decision and selection strategy: The agent makes decisions based on the current state information. The action of the agent is a binary vector, where each element represents whether a client is selected to participate in the training of the current round. The agent selects the best client by evaluating the state of each client to optimize the training process.

[0110] Action space A k Is represented as:

[0111] A k = [a1, a2, …, a Mk

[0112] Where each element a t ∈ {0, 1}, indicating whether to select the t-th client in cluster k (1 for selection, 0 for non-selection).

[0113] Dynamic exploration strategy: To balance exploration and exploitation, the agent adopts a phased ε-greedy decay strategy, exploring in a larger range in the initial stage (ε = 0.5). As the training progresses, the value of ε gradually decays to 0.2 and then stabilizes at 0.1 in the later stage. This strategy can help the agent gradually optimize client selection and avoid falling into local optimal solutions.

[0114] Forced exploration mechanism: For clients that have not participated in the training in the past 10 rounds, the agent will forcibly select these clients with a probability of 30% to avoid biases in model learning caused by their long-term non-participation.

[0115] Frequency limit mechanism: When a certain client participates in the training more than 4 times within the past 5 rounds, the agent will forcibly skip this client to avoid resource consumption problems caused by its frequent participation.

[0116] Among them, the agent decision flow chart is as Figure 3 shown. The agent module uses a Value Decomposition Network (VDN) to achieve collaborative decision-making, and the agents share neural network parameters to reduce the training complexity.

[0117] Step (4), local training and model uploading.

[0118] In each round of training, after being selected by the agent, the client conducts local training and uploads the model updates generated during the training process to the central server.

[0119] ​Local Training: The selected clients perform local training based on the global model. The clients use their local datasets for training and optimize the local model by calculating the updates of the model parameters. Each client only updates the parameters of the model without uploading the original data, ensuring data privacy protection. The local training process of the clients includes standard steps such as forward propagation, error calculation, backpropagation, and gradient update. The efficiency of the training process is affected by the computing power and data volume of the clients. Therefore, the agent selection strategy will preferentially select clients with stronger computing power and richer data distribution.

[0120] Model Upload: After each client completes local training, it uploads the generated model updates (such as weight updates) to the central server. The amount of data uploaded is usually small because it only contains the updates of the model rather than the original data, which greatly reduces the communication burden. To further reduce communication latency and energy consumption, the present invention can use compression techniques (such as model weight compression, sparsification techniques, etc.) to further reduce the size of the uploaded model.

[0121] Step (5), Global Model Aggregation and Feedback.

[0122] Once the clients upload the local model updates, the central server will receive the model updates from all clients and aggregate them to generate the global model.

[0123] Optimize the client selection strategy through a multi-objective reward function. The reward function includes:

[0124]

[0125] Where:

[0126] ΔAcc: Improvement in the test accuracy of the global model;

[0127] H t : The maximum communication and training latency in the current round;

[0128] De k : Difference in client control variables;

[0129] Num(K t ) : The number of clusters covered in the current round;

[0130] ω1, ω2, ω3, β are weight coefficients.

[0131] Model aggregation: The central server uses weighted average or other aggregation methods to merge the model updates uploaded by all clients, thus obtaining a new global model. Usually, the aggregation weights are adjusted according to factors such as the local data volume and computing power of the clients to ensure that the update contributions of different clients are reasonably reflected. The global model aggregated by the server is broadcast to the agents of each cluster for their use in the next round of training.

[0132] Feedback mechanism: After the global model is updated, the agents will adjust the selection strategy based on the feedback of the global model. The update of the global model reflects the effect of the current client selection strategy. The agents optimize the client selection for the next round according to the feedback information such as the test accuracy and communication delay of the global model. The feedback mechanism includes a multi-objective reward function, which considers the following aspects:

[0133] Model accuracy: The improvement of the test accuracy of the global model in the current round.

[0134] Communication delay: The communication delay between the client and the central server in this round of training.

[0135] Gradient deviation: The gradient difference between the local model of the client and the global model.

[0136] Inter-cluster diversity reward: Encourage the selection of clients from different clusters to avoid biases caused by over-training within the cluster.

[0137] The optimization objective is expressed by the following formula:

[0138]

[0139] Where:

[0140] A is the client selection strategy matrix;

[0141] Acc(T) is the test accuracy of the global model in round T;

[0142] H t Is the total processing delay in round t;

[0143] B t Is the total communication energy consumption in round t.

[0144] Step (6), repeated optimization and termination conditions.

[0145] The client selection method of the present invention is an iterative optimization process. After each round of training, the agents will adjust the strategy based on the feedback of the global model and select new clients to participate in the training. Through repeated optimization, the accuracy and convergence speed of the model can be continuously improved.

[0146] Iterative Optimization: The client selection process is repeatedly carried out in multiple training rounds. After each round of training, the agent adjusts its strategy according to the feedback of the new global model, selects new clients to participate in training until the global model reaches the predetermined performance goal.

[0147] Termination Conditions: The training process stops according to preset conditions. For example, when the accuracy of the global model on the test set reaches a certain target value, or after a predetermined maximum number of training rounds, the training process can end.

[0148] The method of the present invention uses the mean shift clustering algorithm, the multi-agent reinforcement learning (MARL) framework and the gradient correction technology to solve the problems of optimizing the efficiency and model performance of federated learning in scenarios such as data heterogeneity (Non-IID), resource constraints of edge devices (unequal computing power, communication delay, energy consumption limitation), etc. Based on multi-dimensional state perception (including gradient deviation quantization, resource load, dynamic behavior and cluster-level data distribution characteristics), combined with dynamic exploration strategies (phased ε-greedy decay, forced selection mechanism) and inter-cluster diversity rewards, intelligent client screening and resource balance are achieved. Without increasing the additional communication load, the convergence speed of the global model is significantly improved, the communication energy consumption is reduced, and the generalization performance of the model on unknown data distributions is enhanced, effectively adapting to the efficient federated learning deployment in the edge computing environment.

[0149] The present invention also provides a federated learning system, including a number of clients and a central server, and the clients are configured to:

[0150] Conduct local probing training based on the current global model, calculate the local model update, and collect the probing training loss, the corrected loss corrected by the control variable, and the probing training delay; record the historical communication delay, training delay, communication energy consumption, local dataset size, participation frequency, recent participation interval and training round index of the client; upload the status data to the central server;

[0151] And when selected by the central server, conduct local training, calculate the model update and upload it to the central server;

[0152] The central server is configured to:

[0153] Based on the local data distribution of the clients, use the mean shift clustering algorithm to divide the clients into multiple data homogeneous clusters, generate cluster identifiers and cluster statistical features; each cluster corresponds to a reinforcement learning agent, which is responsible for dynamically selecting clients within the cluster to participate in training;

[0154] Perform agent state modeling and action decision-making. The agent state space includes detection training loss, corrected loss, control variable difference, historical communication delay, training delay, communication energy consumption, local dataset size, participation frequency, recent participation interval, training round index, cluster identifier, and cluster statistical features. The agent dynamically selects clients based on multi-dimensional state information, and adopts a dynamic exploration strategy and frequency limit mechanism to optimize client selection;

[0155] And aggregate the model updates uploaded by clients, optimize the global model and broadcast it to each cluster; and adjust the strategy according to the feedback of the global model, and optimize the client selection strategy through a multi-objective reward function.

[0156] The above embodiments are only used to illustrate the technical solutions of the present invention and are not intended to limit them. Although the present invention has been described in detail with reference to the above embodiments, those of ordinary skill in the art should understand that: the specific implementation manners of the present invention can still be modified or equivalently replaced, and any modification or equivalent replacement without departing from the spirit and scope of the present invention shall be covered by the protection scope of the claims of the present invention.

Claims

1. A method for selecting a federated learning client based on reinforcement learning, characterized in that, It includes the following steps: The client conducts local probing training and status collection, including: each client conducts local probing training based on the current global model, calculates the local model update, and collects the probing training loss, the corrected loss corrected by the control variable, and the probing training delay; records the client's historical communication delay, training delay, communication energy consumption, local dataset size, participation frequency, recent participation interval, and training round index; The central server clusters and partitions the clients, including: based on the local data distribution of the clients, uses the mean shift clustering algorithm to partition the clients into multiple data homogeneous clusters, generates cluster identifiers and cluster statistical features; each cluster corresponds to a reinforcement learning agent responsible for dynamically selecting clients within the cluster to participate in training; The central server conducts agent state modeling and action decision-making, including: the agent dynamically selects clients based on multi-dimensional state information including the probing training loss, corrected loss, control variable difference, historical communication delay, training delay, communication energy consumption, local dataset size, participation frequency, recent participation interval, training round index, cluster identifier, and cluster statistical features; adopts a dynamic exploration strategy and a frequency limit mechanism to optimize client selection; The client conducts local training and model upload, including: the selected client conducts local training, calculates the model update, and uploads it to the central server; The central server conducts global model aggregation and feedback, including: the central server aggregates the model updates uploaded by the clients, optimizes the global model, and broadcasts it to each cluster; the agent adjusts the strategy according to the feedback of the global model and optimizes the client selection strategy through a multi-objective reward function.

2. The method according to claim 1, characterized in that The mean shift clustering algorithm iteratively updates the clustering center of each client by calculating the mean within the feature space of each client and finally converges to a set of representative clusters; The cluster statistical features include: Data class distribution entropy: measures the distribution diversity of data classes within the cluster; Data volume variance: measures the unevenness of the data volume of clients within the cluster; Client resource mean: calculates the average of the computing power and communication energy consumption of clients within the cluster.

3. The method according to claim 1, characterized in that, The state of the agent is as follows: s t = [l t , l′ t , De t , T comm,t , T train,t , E t , D t , G t , F t , C t , S stats,t ​ Where: l t is the detection training loss for client t; l′ t is the loss after the control variable correction for the client t; De t is the difference between the local control variable of the client and the global model control variable; T comm,t is the communication delay of client t; T train,t is the local training time of client t; E t is the communication energy consumption of client t; D t is the size of the local dataset of the client t; G t is the frequency of client t participating in training; F t is the time interval since the client t last participated in training; C t Is the identifier of the cluster to which the client t belongs; S stats,t It is the statistical feature of the cluster where the client t is located.

4. The method according to claim 3, wherein Differences De between the local control variables of the client and the global model control variables t The calculation method is as follows: De t = ||c t - c||2 Among them, c t is the local control variable of the client t, and c is the global control variable.

5. The method according to claim 1, wherein Assume that each cluster C k contains M k clients, and the action space A of the agent k is expressed as: A k = [a1, a2, …, a Mk ​ where the action of the agent is represented as a binary vector, where each element a t ∈ {0, 1}, indicating whether to select the t-th client in cluster k.

6. The method according to claim 1, characterized in that The action selection strategy includes: Dynamic exploration strategy: ε-greedy decay in stages, where ε is the decay factor; Forced exploration mechanism: for clients that have not participated in training for more than N rounds, select them with a specified probability; Frequency limit: if the number of participations of a client in the past specified rounds ≥ the specified threshold, force it to skip 1 round.

7. The method according to claim 1, wherein Each agent uses a multi-agent reinforcement learning algorithm to select clients. The algorithm updates the strategy based on the multi-objective reward function feedback by the global model. The reward function is as follows: Where: ΔAcc is the improvement in the test accuracy of the global model; H t is the maximum communication and training delay for the current round; De k is the difference in client control variables; Num(K t ) is the number of clusters covered in the current round; ω1, ω2, ω3, β are weight coefficients.

8. The method according to claim 7, wherein The optimization objective is represented by the following formula: Where: A is the client selection strategy matrix; Acc(T) is the test accuracy of the global model at round T; H t is the total processing delay at round t; B t is the total communication energy consumption for round t.

9. A federated learning system, characterized in that, It includes several clients and a central server. The client is configured to: Conduct local probing training based on the current global model, calculate the local model update, and collect the probing training loss, the corrected loss corrected by the control variable, and the probing training delay; Record the historical communication latency, training latency, communication energy consumption, local dataset size, participation frequency, recent participation interval, and training round index of the client; upload the status data to the central server; and when selected by the central server, perform local training, calculate the model update, and upload it to the central server; The central server is configured to: Based on the local data distribution of the clients, use the mean shift clustering algorithm to divide the clients into multiple data homogeneous clusters, and generate cluster identifiers and cluster statistical features; Each cluster corresponds to a reinforcement learning agent, which is responsible for dynamically selecting clients within the cluster to participate in training; Perform agent state modeling and action decision-making. The agent state space includes the detection training loss, corrected loss, control variable difference, historical communication latency, training latency, communication energy consumption, local dataset size, participation frequency, recent participation interval, training round index, cluster identifier, and cluster statistical features. The agent dynamically selects clients based on multi-dimensional state information, and adopts a dynamic exploration strategy and a frequency limit mechanism to optimize client selection; and aggregate the model updates uploaded by the clients, optimize the global model, and broadcast it to each cluster; and adjust the strategy according to the feedback of the global model, and optimize the client selection strategy through a multi-objective reward function.

Citation Information

Cited By

  • Multi-robot collaborative boarding method and system based on reinforcement learning

    CN120469431A

  • Federal learning training method and system for weak network environment

    CN121098743A

  • Model training method and device, electronic equipment, storage medium and program product

    CN121146118A

  • Federal learning-based medical data privacy calculation method and system

    CN121728085A

  • Federal learning client selection method based on Lyapunov optimization in mobile edge computing network

    CN121984878A