Heterogeneous perception traffic prediction method based on federated learning
Through the multi-dimensional personalized federated learning and time window training method, heterogeneity and data quality problems between clients are solved, efficient traffic prediction and privacy protection are achieved, and the robustness and accuracy of the model are improved.
Patent Information
- Application Number
- CN202510447326.7
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-04-10
- Publication Date
- 2025-07-08
AI Technical Summary
The existing traffic prediction method based on federated learning has problems of limited performance and insufficient privacy protection when dealing with the problems of spatial feature heterogeneity, time coverage heterogeneity and data quality among clients.
Multi-dimensional personalized federated learning is used for client clustering, local model training is used to use time windows, global detection and local denoising, and combined with methods of global sharing and personalized parameter aggregation, we solve the problems of data heterogeneity and missing.
Improve the accuracy and privacy protection of traffic forecasts, reduce the burden of data transmission, and enhance the robustness and accuracy of the model.
Smart Images

Figure CN120279709A_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the technical field of traffic prediction, and particularly relates to a heterogeneous perception traffic prediction method based on federated learning. Background Art
[0002] Traffic prediction aims to predict traffic-related data (such as traffic flow, speed, and occupancy rate, etc.) of each region based on historical traffic sequence data. Accurate traffic prediction can promote efficient traffic resource scheduling, help traffic departments and relevant agencies better plan the traffic system, so as to cope with traffic congestion problems in different time periods and regions.
[0003] In view of the wide application of traffic flow prediction, many prediction methods have been proposed by all sectors of society. These methods aim to effectively capture the spatio-temporal relationship between urban regions and have achieved quite high accuracy. Most existing studies tend to design models through centralized learning strategies. Traditional deep learning time series prediction models, such as LSTM (Long Short-Term Memory model) and the current most advanced models in the direction of traffic flow prediction, ST-SSL (Spatio-Temporal Self-Supervised Learning model) and DyHSL (Dynamic Hypergraph Structure Learning model), all only consider centralized learning strategies, that is, aggregating all traffic data within the region to jointly train a model. This strategy requires transmitting a large amount of traffic data from scattered regions to the central server for model training, which has the risk of leaking the privacy of personal and sensitive information in traffic data. Therefore, it becomes crucial to maintain the decentralization of traffic data to reduce communication burden and protect privacy while predicting traffic flow.
[0004] Federated learning is an innovative distributed computing paradigm that solves privacy and efficiency problems by dispersing model training among unconnected clients. Each client corresponds to a region, and these clients use their respective private data to jointly train the model and only exchange intermediate parameters with the server for model aggregation. This method helps to solve the problems of privacy protection and data security, and at the same time can also make full use of region-specific information to improve the accuracy of traffic flow prediction.
[0005] In summary, designing a heterogeneous perception traffic prediction method based on federated learning has become an urgent need in academia and industry. However, there will be the following problems in realizing traffic prediction based on federated learning on the basis of existing technologies:
[0006] Firstly, there is spatial feature heterogeneity among clients. Due to the changes in traffic-related characteristics (such as traffic flow patterns, road structures, point-of-interest distributions, etc.) in different regions, there are inherent differences in the spatial data statistics of different clients. Existing centralized prediction algorithms learn spatial heterogeneity by integrating additional data sets from different regions, but tend to ignore local-specific traffic patterns when applied in the federated scenario, resulting in limited performance.
[0007] Secondly, there is time coverage heterogeneity among clients. Clients in different regions may lack some data due to factors such as different start times and sampling rates of sensor settings, interruptions caused by system updates, and road controls. Although existing algorithms interpolate missing data before model training, they ignore the effectiveness of the global traffic patterns in other regions during the same period on the local missing data and traffic patterns, resulting in limited performance.
[0008] In addition, due to data noise or outliers, the data quality among clients is also uneven. Although some algorithms perform comprehensive data denoising and filtering during model training, it may cause the prediction results to deviate from the actual traffic conditions due to filtering out congestion, accidents, or temporary road controls. Summary of the Invention
[0009] In view of the above, the present invention provides a heterogeneous-aware traffic prediction method based on federated learning, which can provide general federated functions for various centralized prediction models and support traffic flow, speed, and occupancy prediction tasks.
[0010] A heterogeneous-aware traffic prediction method based on federated learning includes the following steps:
[0011] (1) Use clients to obtain traffic data in each region;
[0012] (2) The central server clusters the clients through a multi-dimensional personalized federated learning method according to the spatial feature distribution of the clients;
[0013] (3) The clients use a federated training method based on a time window to perform local model training;
[0014] (4) The central server performs global detection on the models uploaded by the clients. For the detected abnormal models, the clients asynchronously perform local denoising on the local traffic data;
[0015] (5) The central server performs personalized aggregation and global sharing on the models uploaded by the clients according to the clustering results.
[0016] Furthermore, the traffic data consists of a large amount of traffic information, and each piece of information includes traffic data features and their corresponding timestamps. The traffic data features include traffic flow, speed, and occupancy.
[0017] Further, the specific implementation of step (2) is as follows: First, the client embeds and uploads the regional features including the road network structure and attributes, weather conditions, and points of interest in the area as additional data together with the traffic data to the central server after feature encoding. The central server collects the features uploaded by all clients, generates multi-dimensional positive samples through data augmentation, and then clusters the clients through contrastive learning using feature anchors and the corresponding multi-dimensional positive samples. Through the contrastive learning of the client's personalized features, the central server can effectively capture the local traffic patterns, learn the traffic patterns unique to the local area and the globally common traffic patterns during the aggregation process, and maximize the distance between similar clients without the need to determine the number of clusters in advance.
[0018] Further, the specific implementation of step (3) is as follows: First, the client splits the local traffic data into multiple partitions of equal time length and trains the local model (for traffic prediction) in sequence according to the order of the partitions. The model will only move to the next partition to continue training after it converges in one partition until the model is trained in all partitions. This technical feature helps the missing areas learn the global traffic patterns and data periodicity of other areas, which can not only solve the problem of uneven time coverage but also enhance the learning of traffic data periodicity.
[0019] For any partition, if there is traffic data in the partition, the client uses the data in the partition for local model training and gradient calculation, uploads the model to the central server to participate in the federated aggregation, and then uses the aggregated model parameters distributed by the central server to update the local model; if there is no traffic data in the partition, the client uses the data in the adjacent partition for local model training and gradient calculation, and does not upload the model to the central server to participate in the federated aggregation, but uses the aggregated model parameters distributed by the central server to update the local model; enabling the client model to learn the traffic patterns of other areas and the local specific traffic patterns to supplement the missing data.
[0020] Further, the specific implementation of step (4) is as follows: For any client, it uploads the model to the central server. The central server stores the model and compares it with the historical average model of the client to obtain the similarity result. If the similarity result is lower than the set threshold, the central server marks the model as abnormal and notifies the client to perform local denoising on the local traffic data. This technical feature helps to globally detect the uploaded gradients or models during model training to identify and filter out noise, effectively solving data heterogeneity and achieving better performance than other methods that perform data denoising before model training.
[0021] Further, the specific implementation of step (5) is as follows: In each round of training, each client selects the top K parameters with the highest absolute gradient values from the model it uploads, where K is a natural number greater than 1; take the intersection of the parameters selected by all clients, and the parameters in the intersection part are the global parameters, and the parameters in the non-intersection part and the unselected parameters are all personalized parameters; according to the clustering results, the central server aggregates the personalized parameters among clients in the same class, aggregates the global parameters among all clients, and distributes the aggregated personalized parameters to the corresponding clients in categories, and uniformly distributes the aggregated global parameters to all clients.
[0022] A computer device includes a memory and a processor. A computer program is stored in the memory, and the processor is configured to execute the computer program to implement the above-mentioned heterogeneous perception traffic prediction method based on federated learning.
[0023] A computer-readable storage medium stores a computer program, and when the computer program is executed by a processor, it implements the above-mentioned heterogeneous perception traffic prediction method based on federated learning.
[0024] Based on the above technical solutions, the present invention has the following beneficial technical effects:
[0025] 1. The present invention provides a set of heterogeneous perception traffic prediction solutions for traffic patterns based on federated learning. It designs a unified heterogeneous perception framework using federated learning to support existing centralized traffic prediction models.
[0026] 2. The present invention uses multi-dimensional positive sample contrast learning to group clients with similar traffic flow data distributions together, enabling clients in the same class to jointly train the model and avoiding the influence of data heterogeneity between different clients.
[0027] 3. The present invention uses data partitioning to train the models of each stage in sequence based on time windows, reducing the impact of data missing on traffic prediction.
[0028] 4. The present invention uses noise detection for global detection and local denoising to ensure the quality of client data. BRIEF DESCRIPTION OF THE DRAWINGS
[0029] Figure 1 It is a schematic flow chart of the heterogeneous perception traffic prediction method based on federated learning of the present invention. DETAILED DESCRIPTION OF THE EMBODIMENTS
[0030] In order to describe the present invention more specifically, the technical solutions of the present invention will be described in detail below with reference to the drawings and specific embodiments.
[0031] As Figure 1As shown in the figure, the heterogeneous perception traffic prediction method based on federated learning of the present invention is applied to a terminal, and specifically includes the following steps:
[0032] S11: The client obtains traffic data of each region, where each piece of information includes time and corresponding traffic data features, and the traffic data features include flow, speed, and occupancy rate.
[0033] Specifically, in the traffic system, a large amount of traffic data is generated at each moment. By summarizing the traffic data generated by the sensor nodes at fixed time intervals, the traffic characteristics of the node in this time period can be obtained, including flow, speed, and occupancy rate.
[0034] S12: The central server adopts multi-dimensional personalized federated learning to cluster the clients according to the spatial feature distribution of the clients.
[0035] Specifically, the road network structure and attributes, weather conditions, point-of-interest distribution, etc. of each region are used as multi-dimensional additional data to form a set Based on this, an embedded feature set can be obtained which consists of the feature encoding f() on each client. Then the client uploads the feature set to the central server, and the central server generates multi-dimensional similar positive samples through multi-dimensional data augmentation aug(); finally, the central server uses the feature anchor points and the corresponding multi-dimensional positive samples to perform client clustering through contrastive learning, maximizing the distance between similar clients without having to determine the number of clusters in advance.
[0036] S13: Each client uses federated training based on a time window for local model training.
[0037] Consider a client with a dataset D. In the training preparation stage, D is split into partitions of equal time length TW, i.e., D = {d (1) , …, d (TW)}, where TW represents the number of time windows. In the model training stage, the client trains the model sequentially according to the order of the time windows, and the model will only move to the next time window after a sufficient number of rounds or when it converges in the current window; the training process will continue until the model has been trained in all windows. Due to the problem of time coverage deviation, the amount of data available to the client in different time windows is inconsistent, and there may even be no data in the client in some time windows, that is, there is no data available to the client in the t-th time window. In this case, the model training cannot continue. Therefore, the present invention adopts a federated training strategy, and the specific process is as follows:
[0038] If there is data in the t-th time window of client c, that is The client uses data for model training and gradient calculation, uploads the model to the server for federated aggregation, and uses the FedAvg algorithm to aggregate model updates to the local model. The model training and update process is as follows:
[0039]
[0040] Where: w is the client model, is the aggregated model, i is the training round, η is the learning rate, is the gradient of the model w trained on the dataset d (t) and k is the number of useful clients in the t-th time window.
[0041] If there is no data in the t-th time window of client c, i.e., it will use the data closest to the current time window for model training and gradient calculation. To avoid affecting the learning of traffic patterns by other clients in the current time window, this client will not upload the model to the server or participate in federated aggregation. Instead, it will use the aggregated model and the locally trained model during model update, enabling the client model to learn traffic patterns in other regions and local-specific traffic patterns to supplement the missing data.
[0042] First, check the next time window, then the previous time window, and so on, until a time window containing available data is found; even if there is only a small amount of data in the current time window, the module will continue with model training within the current time window. The model training and update process is as follows:
[0043]
[0044] Where: d near(t) is the data of the window closest to the t-th time window.
[0045] S14: After the client uploads its model to the server, the server performs global model detection, and the client performs local denoising asynchronously.
[0046] If an abnormal model is detected, indicating the presence of noise in the local data used for training, the client performs local denoising. Through this fine-grained model detection and data denoising, real noise in the real dataset can be identified and filtered out, thereby improving the accuracy of federated traffic prediction. Specifically:
[0047] After each client uploads its model, the server stores it and compares it with the historical average model to obtain a model similarity result within the range of [-1, 1]. When the similarity is lower than the threshold ρ, the server marks the model as abnormal and notifies the client to perform partial denoising on the local data, as follows:
[0048]
[0049] Wherein: is the historical average model of the client, is the model similarity between model w and the historical average model dr is the detection result of the client model; dr being true indicates an abnormal model, and being false indicates a normal model.
[0050] To minimize the impact of model detection and data denoising on training efficiency, these detection and denoising tasks are postponed in the background, and this delay method helps reduce the occupancy of computing resources during model training.
[0051] S15: The server aggregates all client models based on the client clustering results obtained in step S12.
[0052] This step includes two parts: the server side and the client side:
[0053] On the server side, first collect the data feature distributions of the clients, generate multi-dimensional positive samples, and use these samples for client clustering through contrastive learning. For each time window and training round, initialize the client model detection results and the aggregated model; then, for each cluster, obtain the client models and perform model detection on each model; subsequently, the server aggregates the personalized parameters of the models belonging to the same cluster and the global parameters of all client models; finally, transmit the aggregated models and model detection results to the clients. The model aggregation process on the server side is as follows:
[0054] In each round of training, the client selects the top K parameters with the highest absolute gradient values and takes their intersection positions in the model. These intersection positions represent the global parameters, while the other positions represent the personalized parameters. In the same cluster, the personalized parameters are aggregated as shown in the following formula:
[0055]
[0056] Wherein: is the personalized parameter of client i in cluster j, is the personalized aggregation model of cluster j, |cluster j | is the number of clients in cluster j, n ij is the data volume of client i in cluster j.
[0057] Among all clients, the global parameters are globally aggregated as shown in the following formula:
[0058]
[0059] Wherein: is the global parameter of client i, is the global aggregation model of all clients, |C| is the total number of clients, and n i is the data volume of client i.
[0060] At the client side, first, the model is initialized; then, for each time window and training round, the client obtains the dataset of the current time window. If there is no available data in data partition d (t) the client obtains data d near(t) near the current time window and then obtains the mini-batch data and performs gradient calculation; subsequently, the client obtains the model aggregation of the corresponding time window with available data to update the local model. In contrast, if there is available data in the dataset d (t) of the current time window, the client obtains the mini-batch data, performs gradient calculation and model update, and uploads the updated model to the server. Finally, the client obtains the aggregated model and the model detection result. If the result is true, the client performs data denoising on the batch. The local model training process of the client is as follows:
[0061] S5-1: Start by dividing the dataset D into TW partitions. For each partition d (t) the traffic data is input into the model w of each client as shown in the following formula:
[0062]
[0063] where: is the prediction result of each partition d (t) .
[0064] S5-2: The model is optimized by minimizing the loss of the d (t) dataset as shown in the following formula:
[0065]
[0066] where: |d (t) | is the data volume of partition d (t) , and abs() is the function to calculate the absolute value.
[0067] S5-3: The average loss is obtained by combining the losses of all clients as shown in the following formula:
[0068]
[0069] S5-4: Train the local model using the backpropagation algorithm;
[0070] S5-5: The client uploads the trained model to the server for aggregation;
[0071] S5-6: Repeat steps S5-2 to S5-5 until the loss converges.
[0072] The above description of the embodiments is to enable those of ordinary skill in the art to understand and apply the present invention. It is obvious that those who are familiar with the technology in this field can easily make various modifications to the above embodiments and apply the general principles described herein to other embodiments without creative labor. Therefore, the present invention is not limited to the above embodiments, and all improvements and modifications made by those skilled in the art based on the disclosure of the present invention should fall within the protection scope of the present invention.
Claims
1. A heterogeneous perception traffic prediction method based on federated learning, comprising the following steps: (1) Use the client to obtain traffic data in each region; (2) The central server clusters the clients through a multi-dimensional personalized federated learning method according to the spatial feature distribution of the clients; (3) The client uses a federated training method based on a time window to perform local model training; (4) The central server performs global detection on the models uploaded by the clients. For the detected abnormal models, the clients asynchronously perform local denoising on the local traffic data; (5) The central server performs personalized aggregation and global sharing on the models uploaded by the clients according to the clustering results.
2. The heterogeneous perception traffic prediction method based on federated learning according to claim 1, characterized in that: The traffic data consists of a large amount of traffic information. Each piece of information contains traffic data features and their corresponding timestamps. The traffic data features include traffic flow, speed, and occupancy rate.
3. The heterogeneous perception traffic prediction method based on federated learning according to claim 1, wherein: The specific implementation method of step (2) is as follows: First, the client embeds the regional features including the road network structure and attributes, weather conditions, and points of interest in the region as additional data together with the traffic data after feature encoding and uploads them to the central server. The central server collects the features uploaded by all clients, generates multi-dimensional positive samples through data augmentation, and then uses feature anchors and the corresponding multi-dimensional positive samples to cluster the clients through contrast learning.
4. The heterogeneous perception traffic prediction method based on federated learning according to claim 1, characterized in that: The specific implementation method of step (3) is as follows: First, the client splits the local traffic data into multiple partitions of equal time length, and trains the local model in sequence according to the order of the partitions. The model will move to the next partition to continue training only after it converges in one partition until the model is trained in all partitions; For any partition, if there is traffic data in the partition, the client uses the data in the partition to perform local model training and gradient calculation, and uploads the model to the central server to participate in federated aggregation, and then uses the aggregated model parameters distributed by the central server to update the local model; if there is no traffic data in the partition, the client uses the data in the neighboring partition to perform local model training and gradient calculation, and will not upload the model to the central server to participate in federated aggregation, but will use the aggregated model parameters distributed by the central server to update the local model.
5. The heterogeneous perception traffic prediction method based on federated learning according to claim 1, characterized in that: The specific implementation method of step (4) is as follows: For any client, it uploads the model to the central server. The central server stores the model and compares the model with the historical average model of the client to obtain a similarity result. If the similarity result is lower than the set threshold, the central server marks the model as abnormal and notifies the client to perform local denoising on the local traffic data.
6. The heterogeneous perception traffic prediction method based on federated learning according to claim 1, characterized in that: The specific implementation method of the step (5) is as follows: in each round of training, each client selects the top K parameters with the highest absolute gradient values from the model it uploads, where K is a natural number greater than 1; take the intersection of the parameters selected by all clients, and the parameters in the intersection are the global parameters, and the parameters in the non-intersection part and the unselected parameters are all personalized parameters; according to the clustering results, the central server aggregates the personalized parameters among clients in the same class, aggregates the global parameters among all clients, and distributes the aggregated personalized parameters to the corresponding clients in categories, and uniformly distributes the aggregated global parameters to all clients.
7. A computer device, comprising a memory and a processor, wherein a computer program is stored in the memory, and characterized in that: The processor is used to execute the computer program to implement the heterogeneous perception traffic prediction method based on federated learning according to any one of claims 1 to 6.
8. A computer-readable storage medium storing a computer program, characterized in that: When the computer program is executed by the processor, it implements the heterogeneous perception traffic prediction method based on federated learning according to any one of claims 1 to 6.
Citation Information
Cited By
Federal learning-based dynamic space-time diagram traffic flow prediction method and related device
CN121904994A