Longitudinal graph federal learning method based on multi-head attention mechanism

By using a multi-head attention mechanism to weighted aggregate local node embeddings in vertical graph federated learning, the problem of insufficient expression capabilities of global node embeddings is solved, and higher prediction accuracy is achieved, providing technical support for the construction of smart transportation.

CN120069008AInactive Publication Date: 2025-05-30LANZHOU UNIV
View PDF 4 Cites 0 Cited by

Patent Information

Application Number
CN202510534836.8
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-04-27
Publication Date
2025-05-30
Estimated Expiration
Not applicable · inactive patent

AI Technical Summary

Technical Problem

In the aggregation process of global node embedding, the existing vertical graph federated learning technology lacks feature expression capabilities and cannot effectively capture the complex relationship between local node embedding and global node embedding.

Method used

The multi-head attention mechanism is used to weighted aggregate the local node embedding after differential privacy perturbation, and update the server's global node embedding. This method calculates the attention value between the local node embedding and the global node embedding of each user, performs normalization, and weighted sum based on the attention value to generate a global node embedding with strong feature expression capabilities.

Benefits of technology

The feature expression ability of global node embedding has been improved, the prediction accuracy of tasks such as traffic congestion level prediction, traffic accident prediction, driver driving behavior prediction, and other tasks have been improved, and technical support for the construction of smart cities and smart transportation.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN120069008A_ABST
    Figure CN120069008A_ABST
Patent Text Reader

Abstract

The invention discloses a longitudinal graph federated learning method based on a multi-head attention mechanism, and belongs to the technical field of longitudinal graph federated learning, and the method comprises the steps: obtaining original node features of a traffic monitoring point, carrying out the preprocessing, and enabling a plurality of users to be aligned with the nodes of a server; extracting corresponding node features in each user and the server based on the ordered intersecting node set obtained by alignment, and obtaining an aligned node feature data box; generating local node embedding and performing differential privacy perturbation to obtain perturbed local node embedding; performing weighted aggregation on the disturbed local node embedding through a multi-head attention mechanism, and updating global node embedding of the server; and the server embeds a training classification model through the global nodes to complete node classification of the traffic monitoring points. According to the method, the complex relation between local node embedding and global node embedding can be captured, global node embedding with high feature expression ability is aggregated, and the prediction performance of downstream tasks is improved.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention belongs to the technical field of vertical graph federated learning, and particularly relates to a vertical graph federated learning method based on a multi-head attention mechanism. Background Art

[0002] In the existing technologies that use graph federated learning technology to solve traffic problems, all are scenarios of horizontal graph federation, which are used to solve the problem of insufficient sample data volume of participants. Currently, there is a lack of vertical graph federated technology for solving the problem of insufficient sample features of participants. In vertical graph federated learning technology, some patent methods have been proposed, such as a vertical graph federated recommendation method for two-party users, which mainly applies vertical graph federated technology to commodity recommendation; a vertical graph federated learning method for multiple parties (more than three parties), which uses additive secret sharing to complete the generation of initial node embeddings, uses differential privacy to protect node features, and proposes global node embedding aggregation methods such as mean and concatenation.

[0003] In vertical graph federated learning, multiple users hold different node features, and the server holds node labels. The process includes: (1) users use a local graph neural network model to extract node features and generate local node embeddings (NodeEmbedding); (2) users send the local node embeddings to the server; (3) the server uses methods such as mean and concatenation to aggregate the local node embeddings to obtain global node embeddings; (4) the server completes downstream tasks (such as node classification) according to the global node embeddings; (5) calculate the loss, perform gradient backpropagation, and update the model parameters.

[0004] In this process, the existing technologies have defects in the aggregation of global node embeddings in step (3). Specifically, there are two points: existing technologies, such as the vertical graph federated information recommendation method and related devices based on split learning, only consider vertical graph federation of two parties and do not have technical steps for global node embedding aggregation. Existing technologies, such as the vertical federated learning method and system based on graph neural networks, consider more than three parties, and the aggregation method is too simple, resulting in insufficient feature expression ability of global node embeddings. Specifically, different users hold different features, and the differences in the local node embeddings generated by these users are very large, and their contributions to the global node embeddings are also different. Existing technologies simply average or concatenate the local node embeddings of different users through methods such as averaging and concatenation, without considering the complexity of aggregation and the feature expression ability of global node embeddings at all. Summary of the Invention

[0005] To solve the above technical problems, the present invention proposes a vertical graph federated learning method based on a multi-head attention mechanism to solve the problems existing in the above-mentioned existing technologies.

[0006] To achieve the above object, the present invention provides a vertical graph federated learning method based on a multi-head attention mechanism, including: obtaining the original node features of traffic monitoring points and performing preprocessing; based on the preprocessed original node features, aligning the nodes of multiple parties of users and the server to obtain an ordered intersection node set; based on the ordered intersection node set, extracting the corresponding node features in each user and the server to obtain an aligned node feature data frame; generating a local node embedding based on the aligned node feature data frame and performing differential privacy perturbation to obtain a perturbed local node embedding; performing weighted aggregation on the perturbed local node embedding through a multi-head attention mechanism to update the global node embedding of the server; and the server training a classification model through the global node embedding to complete the node classification of traffic monitoring points.

[0007] Optionally, the preprocessing process includes: converting the original node features of traffic monitoring points into a data frame format to obtain a node feature data frame, where the node unique identifier is the row index and the node features are the column index; cleaning the data frame, removing null values and abnormal nodes, and then filling in the missing data with linear interpolation and normalizing it using the z-score method.

[0008] Optionally, the process of obtaining the ordered intersection node set includes: extracting the unique identifiers of the nodes from the node feature data frames of each user and the server to form corresponding node sets; calculating the intersection of the node sets through private set intersection technology to obtain an ordered intersection node set.

[0009] Optionally, the process of generating a local node embedding includes: performing a non-linear transformation on the original traffic node features through a fully connected layer to generate a hidden representation; inputting the hidden representation into a graph neural network model to output a local node embedding; where the process of generating the hidden representation is as follows.

[0010] 。

[0011] Where represents the aligned node feature data frame, with a total of N nodes, each node having F features, and represent learnable parameters, represents an activation function, taking sigmoid, is the hidden representation.

[0012] Optionally, the process of obtaining the perturbed local node embedding includes: perturbing the local node embedding through the Gaussian mechanism, and for each node embedding , adding the following noise.

[0013] 。

[0014] Among them represents the perturbed node embedding, is an independent and identically distributed random variable sampled from , , and are privacy parameters, represents a normal distribution with a mean of 0 and a variance of .

[0015] Optionally, the process of updating the global node embedding of the server includes: performing a non-linear transformation on the perturbed local node embedding of each user; then, calculating the attention value between the local node embedding of each user and the global node embedding and normalizing it; performing a weighted sum on the local node embeddings of different users according to the normalized attention value to obtain the global node embedding on each head; concatenating the global node embeddings on each head to obtain the global node embedding, and updating the global node embedding of the server.

[0016] Optionally, the classification model is a multi-layer perceptron; when training the classification model, calculate the cross-entropy loss and backpropagate the gradient to the user side, and iteratively update the classification model until convergence.

[0017] The present invention also provides a computer device, including: a memory, a processor, and a computer program stored on the memory and executable on the processor, and the processor executes the computer program to implement the steps of the above method.

[0018] The present invention also provides a computer-readable storage medium, on which a computer program is stored, and when the computer program is executed by a processor, the steps of the above method are implemented.

[0019] The present invention also provides a computer program product, including a computer program, and when the computer program is executed by a processor, the steps of the above method are implemented.

[0020] Compared with the prior art, the present invention has the following advantages and technical effects.

[0021] The present invention provides a vertical graph federated learning method based on a multi-head attention mechanism. First, after obtaining the original node features of traffic monitoring points and preprocessing them, the nodes of multiple parties of users and the server are aligned to obtain an ordered intersection node set. Then, based on the ordered intersection node set, the corresponding node features in each user and the server are extracted to obtain an aligned node feature data frame. Next, local node embeddings are generated based on the aligned node feature data frame and differential privacy perturbation is performed to obtain perturbed local node embeddings. Then, the perturbed local node embeddings are weighted and aggregated through the multi-head attention mechanism to update the global node embeddings of the server. Finally, the server trains a classification model through the global node embeddings to complete the node classification of traffic monitoring points.

[0022] The multi-head attention aggregation method proposed by the present invention can capture the complex relationship between local node embeddings and global node embeddings, aggregate global node embeddings with strong feature expression ability, and thus improve the prediction accuracy of traffic scenario-related tasks such as traffic congestion level prediction, traffic accident prediction, and driver driving behavior prediction. It provides technical support for the construction of smart cities and smart transportation. BRIEF DESCRIPTION OF THE DRAWINGS

[0023] The drawings constituting a part of this application are used to provide a further understanding of this application. The schematic embodiments of this application and their descriptions are used to explain this application and do not constitute an improper limitation to this application. In the drawings.

[0024] Figure 1 It is a training and inference flow chart of the vertical graph federated learning method based on the multi-head attention mechanism according to the embodiment of the present invention.

[0025] Figure 2 It is a schematic diagram of the aggregation method based on the multi-head attention mechanism according to the embodiment of the present invention. DETAILED DESCRIPTION OF THE EMBODIMENTS

[0026] It should be noted that, without conflict, the embodiments in this application and the features in the embodiments can be combined with each other. The following will refer to the drawings and combine the embodiments to detail this application.

[0027] It should be noted that the steps shown in the flowchart of the drawings can be executed in a computer system such as a set of computer executable instructions, and although the logical order is shown in the flowchart, in some cases, the steps shown or described can be executed in a different order than here.

[0028] Embodiment 1: When modeling distributed graph data in traffic scenarios, traffic elements in a city constitute a graph network, and graph nodes represent traffic monitoring points, such as sensors, traffic lights, toll booths, bus stops, etc. on the road. Edges represent roads, which are usually constructed based on traffic topology, such as road connection relationships, vehicle driving trajectories, etc. Node features are in the hands of different institutions. Taking the city traffic management center, map navigation company, and taxi / online car-hailing company as examples, their different feature data are shown in Table 1.

[0029] Table 1 .

[0030] Among them, map navigation companies and online car-hailing companies may have multiple organizations. Those that master node labels and perform downstream task calculations are called servers, and the rest are called users.

[0031] like Figure 1 As shown, this embodiment provides a longitudinal graph federated learning method based on a multi-head attention mechanism, including: obtaining the original node features of traffic monitoring points and preprocessing them; based on the preprocessed original node features, aligning the nodes of multiple users and servers to obtain an ordered set of intersecting nodes.

[0032] As a specific implementation method, the preprocessing process includes: converting the original node features of the traffic monitoring points into a data frame format to obtain a node feature data frame, wherein the node unique identifier is a row index and the node feature is a column index; cleaning the data frame, removing null values ​​and abnormal nodes, filling the missing data with linear interpolation and standardizing it using the z-score method.

[0033] As a specific implementation method, the process of obtaining an ordered set of intersecting nodes includes: extracting the unique identifier of the node from the node feature data frame of each user and server to form a corresponding node set; calculating the intersection of the node set through the privacy set intersection technology to obtain an ordered set of intersecting nodes.

[0034] Specifically, first, the feature information of the node is processed into a data frame format, with the node unique identifier as the row index and the node feature as the column index. Secondly, the data frame is cleaned to remove nodes where all data are null values ​​and data anomalies, and the linear interpolation method is used to fill in the data of nodes with some missing data. Finally, the z-score method is used to standardize the data so that the data is in the range of [-1,1].

[0035] Before the training starts, it is necessary to use the PSI (Private Set Intersection) technology to align all the nodes of the users and the server to obtain an ordered set of intersecting nodes. The data used in the training process is these intersecting nodes.

[0036] Based on the ordered set of intersecting nodes, extract the corresponding node features in each user and server to obtain an aligned node feature data frame.

[0037] Generate local node embeddings based on the aligned node feature data frame and perform differential privacy perturbation to obtain perturbed local node embeddings.

[0038] Specifically, first use a fully connected layer to perform feature transformation on the original node features to generate a hidden representation : .

[0039] Among them represents the original node features. There are a total of N nodes, and each node has F features. and represent learnable parameters. represents the activation function, usually taking sigmoid. is the hidden representation.

[0040] Next, send the hidden representation into a graph neural network model. This graph neural network model can be GCN, GraphSage, GAT, etc. Represent the graph neural network model as , then: .

[0041] Among them represents the local node embedding generated by user , represents the local adjacency matrix of user .

[0042] The local node embedding may leak the user's data privacy, and differential privacy technology can protect this. Differential privacy adds noise to the local node embedding, making it impossible for attackers to obtain exact information about the original data. Specifically, this method uses the Gaussian mechanism to perturb the local node embedding. For each node embedding , add the following noise to it: .

[0043] Among them represents the perturbed node embedding, is an independent and identically distributed random variable sampled from , , and These are all privacy parameters. The mean is 0 and the variance is The normal distribution of . The smaller it is, the stronger the differential privacy protection is, but the prediction accuracy of the model will decrease. Generally, a smaller number is taken, such as 1e-5.

[0044] The perturbed local node embeddings are weighted aggregated through a multi-head attention mechanism to update the server’s global node embeddings.

[0045] As a specific implementation method, the process of updating the global node embedding of the server includes: performing a nonlinear transformation on the perturbed local node embedding of each user; then, calculating the attention value between the local node embedding of each user and the global node embedding and normalizing it; performing a weighted summation of the local node embeddings of different users according to the normalized attention value to obtain the global node embedding on each head; splicing the global node embedding on each head to obtain the global node embedding, and updating the global node embedding of the server.

[0046] Specifically, Figure 2 As shown. First, define the nonlinear transformation: .

[0047] in Indicates input, Represents a nonlinear transformation of the input to the output.

[0048] This embodiment needs to calculate the attention value between the local node embedding and the global node embedding. However, there is no global node embedding at this time, so this embodiment allows the server to maintain a global node embedding , at the beginning of training Initialize randomly and update after each round of aggregation .make represents the number of attention heads, for each head , the attention value between the local node embedding and the global node embedding is calculated as follows.

[0049] .

[0050] in, Indicates The first The global node embedding and user The attention value between the local node embeddings of represents the inner product, function The input is of dimension The output dimension is , the dimension of a single head is . The above formula represents the th head, for the local node embedding and the global node embedding of the th node, first perform a non-linear transformation, and then do matrix multiplication to obtain the attention matrix.

[0051] Subsequently, normalize the attention values: .

[0052] Among them represents the number of users.

[0053] Finally, first perform a non-linear transformation on the local node embeddings of different users, then perform weighted summation according to the attention weights, and finally concatenate the results of

[0054] heads together.

[0055] Among them, is the global node embedding of the th node, represents the concatenation operation.

[0056] After completing the aggregation, update : .

[0057] The server trains a classification model through the global node embedding to complete the node classification of traffic monitoring points.

[0058] As a specific implementation, the classification model is a multi-layer perceptron; when training the classification model, calculate the cross-entropy loss and backpropagate the gradient to the user side, and iteratively update the classification model until convergence.

[0059] Specifically, when the server completes the aggregation of the global node embedding, perform downstream task calculations based on the global node embedding. Examples of classification tasks in traffic scenarios are shown in Table 2.

[0060] Table 2 .

[0061] The specific implementation method is: the server maintains a classifier model , and this classifier model is generally a multi-layer perceptron. Use the global node embedding as the input , and obtain the classification result : .

[0062] Among them, represents The classification results of the nodes, is the number of categories. The classification result of each node is a vector of length indicating the probability that the node belongs to each category.

[0063] During inference, all the calculations are completed at this step, and the predicted category of each graph node is obtained. During training, the loss function and gradients also need to be calculated, and the model parameters are updated.

[0064] The loss function uses the cross-entropy function: .

[0065] Among them, represents the true label of the graph node, is the loss value, is an indicator variable. If the rd node belongs to the th category, the value is 1, otherwise it is 0. represents the probability that the th node belongs to the th category.

[0066] The server takes the derivative of the model with respect to the loss to obtain the gradient , and uses the gradient descent method to update the model parameters. In addition, the server also calculates the gradient of the loss with respect to the local node embedding and sends it to the user . The user calculates the gradient of the model through gradient backpropagation and updates the model parameters.

[0067] So far, one round of training steps is completed. The training of vertical graph federation needs to be iterated for multiple rounds until the model converges.

[0068] This embodiment also provides a computer device, including: a memory, a processor, and a computer program stored on the memory and executable on the processor. The processor executes the computer program to implement the steps of the above method.

[0069] This embodiment also provides a computer-readable storage medium, on which a computer program is stored. When the computer program is executed by a processor, it implements the steps of the above method.

[0070] This embodiment also provides a computer program product, including a computer program. When the computer program is executed by a processor, it implements the steps of the above method.

[0071] The above are only the preferred specific embodiments of the present application, but the protection scope of the present application is not limited thereto. Any changes or substitutions that can be easily conceived by those skilled in the art within the technical scope disclosed in the present application should be covered within the protection scope of the present application. Therefore, the protection scope of the present application should be subject to the protection scope of the claims.

Claims

1. A vertical graph federated learning method based on a multi-head attention mechanism, characterized in that: The following steps are involved: Obtain the original node features of traffic monitoring points and perform preprocessing; Based on the preprocessed original node features, the nodes of multiple users and servers are aligned to obtain an ordered set of intersecting nodes; Based on the ordered intersection node set, the corresponding node features of each user and server are extracted to obtain the aligned node feature data frame; Generate local node embedding based on the aligned node feature data frame and perform differential privacy perturbation to obtain the perturbed local node embedding; The perturbed local node embeddings are weighted aggregated through a multi-head attention mechanism to update the server's global node embeddings. The server completes the node classification of traffic monitoring points by embedding the training classification model through the global node.

2. The method for longitudinal graph federated learning based on multi-head attention mechanism according to claim 1, characterized in that: The preprocessing process includes: The original node features of the traffic monitoring points are converted into a data frame format to obtain a node feature data frame, in which the node unique identifier is the row index and the node feature is the column index. The data frame is cleaned, and after removing null values ​​and abnormal nodes, the missing data is filled with linear interpolation and standardized using the z-score method.

3. The method for longitudinal graph federated learning based on multi-head attention mechanism according to claim 1, characterized in that: The process of obtaining an ordered set of intersection nodes includes: The unique identifier of the node is extracted from the node feature data frame of each user and server to form a corresponding node set; the intersection of the node set is calculated through the privacy set intersection technology to obtain an ordered intersection node set.

4. The method for longitudinal graph federated learning based on multi-head attention mechanism according to claim 1, characterized in that: The process of generating local node embeddings includes: The original traffic node features are transformed nonlinearly through the fully connected layer to generate hidden representations; the hidden representations are input into the graph neural network model and the local node embedding is output; where the hidden representation is generated The process is as follows: ; in Represents the aligned node feature data frame, with a total of N nodes, each node has F features, and represents the learnable parameters, Represents the activation function, taking sigmoid, It is a hidden representation.

5. The method for longitudinal graph federated learning based on multi-head attention mechanism according to claim 1, characterized in that: The process of obtaining the perturbed local node embedding includes: perturbing the local node embedding through the Gaussian mechanism, and for each node embedding , add the following noise: ; in represents the node embedding after perturbation, is from independent and identically distributed random variables sampled from , , and is the privacy parameter, The mean is 0 and the variance is The normal distribution of .

6. The method for longitudinal graph federated learning based on multi-head attention mechanism according to claim 1, characterized in that: The process of updating the server's global node embeddings includes: Perform a nonlinear transformation on the perturbed local node embedding of each user; then, calculate the attention value between the local node embedding of each user and the global node embedding and normalize them; perform weighted summation of the local node embeddings of different users according to the normalized attention value to obtain the global node embedding on each head; concatenate the global node embeddings on each head to obtain the global node embedding, and update the global node embedding of the server.

7. The method for longitudinal graph federated learning based on multi-head attention mechanism according to claim 1, characterized in that: The classification model is a multi-layer perceptron; when training the classification model, the cross entropy loss is calculated and the gradient is back-propagated to the user end, and the classification model is iteratively updated until convergence.

8. A computer device comprising: A memory, a processor, and a computer program stored in the memory and executable on the processor, wherein the processor executes the computer program to implement the steps of the method according to any one of claims 1 to 7.

9. A computer-readable storage medium having a computer program stored thereon, characterized in that: When the computer program is executed by a processor, the steps of the method according to any one of claims 1 to 7 are implemented.

10. A computer program product, comprising a computer program, characterized in that When the computer program is executed by a processor, the steps of the method according to any one of claims 1 to 7 are implemented.

Citation Information

Patent Citations

  • Personalized federal learning method based on multi-head attention mechanism

    CN113378243A

  • Longitudinal federal learning method and system based on graph neural network

    CN118036651A

  • Point cloud data classification method for federal few-sample learning based on privacy protection

    CN118918448A

  • Horizontal and vertical federated learning combined algorithm

    WO2024060410A1