Cross-Domain Graph Model Pretraining Method and System Based on Learnable Graph Patches

By extracting and encoding key graph patches in graph data based on learnable graph patches, and pre-training with multi-task learning, the problem of limited performance of graph neural networks in cross-domain tasks is solved, and the efficient migration performance of graph models in multi-domain tasks is achieved.

CN119782822BActive Publication Date: 2025-05-30ZHEJIANG UNIV
View PDF 2 Cites 0 Cited by

Patent Information

Application Number
CN202510260282.7
Authority / Receiving Office
CN · China
Patent Type
Patents(China)
Current Assignee / Owner
Filing Date
2025-03-06
Publication Date
2025-05-30
Estimated Expiration
2045-03-06

AI Technical Summary

Technical Problem

Existing graph neural networks have limited performance in cross-domain tasks, are difficult to adapt to complex and changeable practical scenarios, and are difficult to achieve efficient representation migration while retaining the structured characteristics of graph data.

Method used

A cross-domain graph model pre-training method based on learnable graph patches is proposed. Through the graph patch extraction module, a single-channel graph patch encoding module and a graph patch aggregation module, a general key graph patch patch is constructed, and pre-trained through feature masking recovery tasks and graph neighbor context prediction tasks, improving the migration performance of graph models.

Benefits of technology

Efficient pre-training of cross-domain graph structure data is realized, which improves the transfer performance of graph models in multi-domain tasks and significantly reduces learning disabilities caused by domain differences.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN119782822B_ABST
    Figure CN119782822B_ABST
Patent Text Reader

Abstract

The present invention discloses a pre-training method and system for cross-domain graph models based on learnable graph patches, belonging to the technical field of graph model training. Obtain pre-training graph datasets of multiple source domains; split the node features of each graph data in the pre-training graph datasets into node token sets of each channel, and the node token sets of each channel and their corresponding graph structures of the channel are used as the key graph patches of the channel; encode and aggregate the key graph patches of each channel in the graph data as the graph data representation; based on the graph data representation result, introduce a feature masking and restoration task and a graph neighbor context prediction task to pre-train each module, and the pre-trained module combination is used as a cross-domain graph model for extracting the target domain graph data representation to implement the downstream tasks of the target domain graph data. The present invention realizes the extraction of transferable information from different source domain graph data by the decomposed key graph patches without auxiliary information, and improves the transfer performance of the graph model in multi-domain tasks.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the technical field of graph model training, and in particular, to a cross-domain graph model pre-training method and system based on learnable graph patches. Background Art

[0002] In the era of data-driven artificial intelligence, graph data, as a common form of unstructured data, is widely used in fields such as social networks, paper citations, and molecular modeling. However, there are significant differences in feature distributions, graph structures, etc. among graph data in different fields, resulting in limited performance of existing graph neural networks (GNNs) in cross-domain tasks. Therefore, how to design an efficient and transferable graph representation learning method to enhance the cross-domain adaptability of graph models is an important research direction in the current graph learning field.

[0003] Existing research mainly relies on domain alignment or generative methods, but these methods usually have too strong assumptions about data distributions and are difficult to adapt to complex and changing actual scenarios. In addition, how to achieve efficient representation transfer while retaining the structural characteristics of graph data is still a technical challenge. Therefore, there is an urgent need for a new type of cross-domain graph learning method that can flexibly extract domain-agnostic general graph representations and support efficient transfer of multi-domain tasks. Summary of the Invention

[0004] To solve the problem that the graph model in the prior art has limited performance in cross-domain transfer tasks, the present invention proposes a cross-domain graph model pre-training method and system based on learnable graph patches. The cross-domain graph model of the present invention consists of a graph patch extraction module, a single-channel graph patch encoding module, and a graph patch aggregation module. By introducing a graph patch partitioning mechanism, a node feature disassembling method, and a graph structure learning method based on the attention mechanism through the learnable graph patch extraction module, general key graph patches that can adapt to multiple domains are constructed. Combining the learnable single-channel graph patch encoding module and the graph patch aggregation module to encode and aggregate the key graph patches to obtain graph data representations, the transfer performance of the graph model in multi-domain tasks is improved.

[0005] To achieve the above object, the technical solution adopted by the present invention is as follows:

[0006] In the first aspect, the present invention proposes a cross-domain graph model pre-training method based on learnable graph patches, including:

[0007] (1) Obtain pre-training graph data sets of multiple source domains, including two or more of citation graph data, social network graph data, molecular graph data, financial network graph data, power network graph data, recommendation system graph data, and traffic network graph data. Each graph data consists of a node feature matrix and an original adjacency matrix;

[0008] (2) Use the graph patch extraction module to split the node features of each graph data in the pre-trained graph dataset into a set of node tokens for each channel, and establish a graph structure based on the set of node tokens for each channel. The set of node tokens for each channel and its corresponding graph structure for that channel are used as the key graph patches for that channel;

[0009] (3) Use the single-channel graph patch encoding module to encode the key graph patches for each channel in the graph data, and then use the graph patch aggregation module to aggregate the encoding results of the key graph patches for all channels as the graph data representation;

[0010] (4) Based on the graph data representation result, introduce the feature masking restoration task and the graph neighbor context prediction task to pre-train the graph patch extraction module, the single-channel graph patch encoding module, and the graph patch aggregation module. The pre-trained module combination is used as a cross-domain graph model for extracting the target domain graph data representation to achieve the downstream tasks of the target domain graph data.

[0011] Preferably, as the present invention, the calculation process of the graph patch extraction module includes:

[0012] (2-1) Split the node features of each graph data in a sliding window manner. The node features are split into node tokens, and the th node token obtained by sequential splitting is used as the node token for the th channel. Each node token has the same dimension;

[0013] (2-2) Construct the corresponding graph structure for each channel based on the set of node tokens for each channel:

[0014] ;

[0015] Among them, represents a neural network model for learning the graph structure, represents the original adjacency matrix of the graph data, represents the set of node tokens for the th channel corresponding to all nodes in the graph data, represents the number of nodes in the graph data, represents the adjacency matrix of the graph structure corresponding to the set of node tokens for the th channel generated by for representing a graph structure; represents the set of node tokens for the

[0016] (2-3) Obtain the key graph patches for each channel according to the set of node tokens for each channel and its corresponding graph structure for that channel:

[0017] ;

[0018] Among them, Indicates the key graph patch of the th channel.

[0019] As a preference of the present invention, the neural network model for learning the graph structure is a neural network model based on the attention mechanism, and the calculation process includes:

[0020] Calculate the attention scores between the node tokens of all nodes corresponding to the th channel in the graph data pairwise;

[0021] Sparsify the attention scores through a non-linear function, and the pairwise sparse attention scores of all nodes in the graph data form the sparse attention adjacency matrix of the th channel of the graph data;

[0022] Perform a residual connection on the sparse attention adjacency matrix of the th channel of the graph data and the original adjacency matrix to form the adjacency matrix corresponding to the graph structure of the node token set of the th channel.

[0023] As a preference of the present invention, the calculation process of the single-channel graph patch encoding module includes:

[0024] (3-1) Use the node token set of the th channel and the original adjacency matrix as the input of the graph neural network to generate the graph patch representation of the th channel based on the original adjacency matrix ;

[0025] And, use the node token set of the th channel and the adjacency matrix corresponding to the graph structure of the node token set of the th channel as the input of the graph neural network to generate the graph patch representation of the th channel based on the new adjacency matrix ;

[0026] (3-2) Fuse the two graph patch representations to obtain the single-channel graph patch encoding, and after traversing all channels, obtain the graph patch encodings of all channels.

[0027] As a preference of the present invention, the process of fusing the two graph patch representations is as follows:

[0028] ;

[0029] Among them, represents a multi-layer perceptron, represents concatenation, represents a non-linear function with a value range of (0,1), represents the Graph patch encoding for one channel.

[0030] Preferably, in the present invention, the graph patch aggregation module adopts a Transformer encoder, and obtains a graph data representation by aggregating the graph patch encodings of all channels.

[0031] Preferably, in the present invention, the feature masking and restoration task includes:

[0032] Masking partial node features of the graph data, and generating a graph data representation according to the masked graph data;

[0033] Using a first decoder to restore the masked node features according to the graph data representation;

[0034] Calculating a feature masking and restoration loss by combining the true masking result and the restoration result.

[0035] Preferably, in the present invention, the graph neighbor context prediction task includes:

[0036] Generating a graph data representation according to the original graph data;

[0037] Using a second decoder to predict the mean feature vector of the p-th order neighbor nodes of each node in the graph data according to the graph data representation, where p = 1, 2,..., P; P represents the set maximum order;

[0038] Calculating a graph neighbor context prediction loss by combining the true result and the prediction result.

[0039] Preferably, in the present invention, the target domain graph data is molecular graph data, and the downstream task of the target domain graph data is a molecular graph classification task.

[0040] In a second aspect, the present invention proposes a cross-domain graph model pre-training system based on learnable graph patches for the above-mentioned cross-domain graph model pre-training method based on learnable graph patches.

[0041] The beneficial effects of the present invention are as follows:

[0042] (1) The present invention realizes a pre-trained graph model for cross-domain graph-structured data. By using the graph patch extraction module, learnable key graph patches are obtained, and the original graph data is split into small-scale and domain-agnostic graph patches. Combining the single-channel graph patch encoding module and the graph patch aggregation module to encode and aggregate the key graph patches to obtain a graph data representation, effectively utilizing the graph-structured data in multiple source domains.

[0043] (2) The graph patch extraction module of the present invention efficiently realizes the mining of common information in graph data between different domains by splitting node features into node tokens, so as to more accurately and reasonably complete the downstream tasks of the target domain graph data.

[0044] (3) Based on the graph data characterization results, the present invention introduces a feature masking and restoration task and a graph neighbor context prediction task to pre-train the graph patch extraction module, the single-channel graph patch encoding module, and the graph patch aggregation module, realizing the extraction of transferable information from different source domain graph data for the decomposed key graph patches without auxiliary information, and improving the transfer performance of the graph model in multi-domain tasks. BRIEF DESCRIPTION OF THE DRAWINGS

[0045] Figure 1 is the overall block diagram of a cross-domain graph model pre-training method based on learnable graph patches shown in an embodiment of the present invention.

[0046] Figure 2 is the flowchart of a cross-domain graph model pre-training method based on learnable graph patches shown in an embodiment of the present invention.

[0047] Figure 3 is a schematic diagram of the calculation process of the graph patch extraction module shown in an embodiment of the present invention.

[0048] Figure 4 is a schematic diagram of the graph patch division method shown in an embodiment of the present invention. DETAILED DESCRIPTION OF THE EMBODIMENTS

[0049] The present invention will be further described and explained below in conjunction with the specific embodiments. The embodiments are only illustrative of the present disclosure and do not delimit the scope of limitation. The technical features of each embodiment of the present invention can be combined correspondingly without conflict.

[0050] The drawings are only schematic diagrams of the present invention and are not necessarily drawn to scale. Some of the block diagrams shown in the drawings are functional entities and do not necessarily correspond to physically or logically independent entities. These functional entities can be implemented in software form, or in one or more hardware modules or integrated circuits, or in different networks and / or processor devices and / or microcontroller devices.

[0051] The flowcharts shown in the drawings are only exemplary descriptions and do not necessarily include all steps. For example, some steps can be decomposed, while some steps can be combined or partially combined. Therefore, the actual execution order may be changed according to the actual situation.

[0052] The core idea of the present invention is to divide the graph data into multiple small-scale and domain-agnostic key graph patches, and optimize the representation ability of the graph patches through joint pre-training. See Figure 1The overall block diagram of the pre-training method for cross-domain graph models based on learnable graph patches. In the present invention, pre-training graph datasets from multiple source domains are collected, and each graph data consists of a node feature matrix and an original adjacency matrix; through the graph patch extraction module, the node features are split into node tokens of a unified dimension, and corresponding graph structures are constructed based on the extracted node tokens to capture the efficient interaction relationships between nodes, obtaining the key graph patches of each channel; the single-channel graph patch encoding module is used to encode the key graph patches of each channel in the graph data, and then the graph patch aggregation module aggregates the encoding results of the key graph patches of all channels as the graph data representation; through multi-task learning methods (feature masking and restoration task and graph neighbor context prediction task), the learnable parameters in the model are jointly trained to enable it to adapt to tasks in different domains.

[0053] As Figure 2 shown, the pre-training method for cross-domain graph models based on learnable graph patches proposed by the present invention mainly includes the following steps:

[0054] S1. Obtain pre-training graph datasets from multiple source domains. Each graph data consists of a node feature matrix and an original adjacency matrix, where the node features represent the attribute information of each node, and the original adjacency matrix represents the connection relationship between nodes.

[0055] The pre-training graph datasets from multiple source domains include at least two domains. For example, a graph dataset jointly composed of a citation graph network and a social graph network, where:

[0056] In the graph data of the citation graph network, the nodes represent papers, the node features are the paper information, and the edges represent the citation relationships between papers; here, the node feature matrix is the matrix composed of paper information, and the original adjacency matrix is the citation relationship matrix between papers.

[0057] In the graph data of the social graph network, the nodes represent users, the node features are the user information, and the edges represent the interaction relationships between users; here, the node feature matrix is the matrix composed of user information, and the original adjacency matrix is the interaction relationship matrix between users.

[0058] In addition, the source domain can also include other domains, such as:

[0059] In the graph data of the molecular graph network, the nodes represent atoms, the node features are the elements of the atoms, and the edges represent the connection relationships between atoms; here, the node feature matrix is the matrix composed of the elements of the atoms, and the original adjacency matrix is the connection relationship matrix between atoms.

[0060] In the graph data of a financial network, nodes represent accounts, node features are account information, and edges represent the fund flow relationship between accounts; here, the node feature matrix is a matrix composed of account information, and the original adjacency matrix is the fund flow relationship matrix between accounts.

[0061] In the graph data of a power grid, nodes represent substations, node features are substation specifications, and edges represent transmission lines; here, the node feature matrix is a matrix composed of substation specifications, and the original adjacency matrix is the transmission line connection matrix.

[0062] In the graph data of a recommendation system, nodes represent users and products, node features are user information and product information, and edges represent the interaction relationship between users and products (such as purchase, evaluation); here, the node feature matrix is a matrix composed of user information and product information, and the original adjacency matrix is the interaction relationship matrix between users and products.

[0063] In the graph data of a transportation network, nodes represent transportation nodes (such as stations, intersections), node features are the geographical information of transportation nodes, and edges represent roads; here, the node feature matrix is a matrix composed of the geographical information of transportation nodes, and the original adjacency matrix is the road connection matrix.

[0064] S2. Use the graph patch extraction module to split the node features of each graph data in the pre-trained graph dataset into a set of node tokens for each channel, and build a graph structure based on the set of node tokens for each channel. The set of node tokens for each channel and its corresponding graph structure for that channel are used as the key graph patches for that channel.

[0065] In this step, as Figure 3 shown, an optional implementation is as follows:

[0066] S21. First, use the graph patch partitioning method to partition each graph data in the pre-trained graph dataset; split the node features of each graph data in a sliding window manner. The node features are split into node tokens. The th node token obtained by sequential splitting is used as the node token for the th channel, ensuring that the dimensions (M and K below) of the node tokens of each graph data are the same. Therefore, the original node features of different graph data will be disassembled into different node tokens, but all node tokens are in the same latent space, which is expressed as:

[0067]

[0068]

[0069]

[0070] Among them, Represents the th node in the graph data, represents the th channel; is the length dimension of the node token, is the sliding window step size, where , represents floor division, represents the dimension of the original node features; since the node feature dimensions of graph data in different fields may vary, adjusting the sliding window step size can ensure that the dimensions of the node tokens of each graph data are consistent. This decomposition method can effectively avoid the problem of inconsistent distributions caused by the direct participation of original features in cross - field learning and lay a foundation for subsequent unified processing. Represents the node token of the th node in the th channel in the data graph, which is a smaller information unit in the node features. Represents the node feature matrix of the graph data, represents the number of node tokens for each node in the data graph, , are respectively the start feature position and the end feature position of the node tokens in the th channel. Taking the node 1 in Figure 4 as an example, the node feature is . Setting the sliding window length to 3 and the step size to 1, node 1 is divided into three node tokens, where , , .

[0071] Traverse all nodes in the graph data, and the set of node tokens of all nodes corresponding to the th channel in the graph data is denoted as .

[0072] S22, and then construct the graph structure of the corresponding channel based on the node tokens of different channels decomposed. The node tokens of each channel will correspond to an independent set of graph structures, which is expressed as:

[0073]

[0074] Among them, is the neural network model used to learn the graph structure, represents the adjacency matrix of the graph data, is the adjacency matrix of the corresponding graph structure generated by for the node tokens of the th channel, is the number of nodes in the graph data. In this embodiment, It can be implemented using a variety of neural network models, such as CNN or neural networks based on the attention mechanism.

[0075] The present invention adaptively learns a dynamic graph structure according to the similarity between node tokens. By calculating the interaction relationship between each pair of node tokens in the latent space, a sparse dynamic adjacency matrix is constructed, and each node token only establishes connections with a number of the most relevant tokens. This not only improves the computational efficiency but also enhances the generalization ability of the model.

[0076] In a specific implementation of the present invention, taking the neural network based on the attention mechanism as an example, the computation process of the graph structure includes:

[0077] First, calculate the attention scores between the node tokens corresponding to all nodes in the graph data for the th channel:

[0078]

[0079] Among them, represents learnable parameters, represents dot product, represents a multi-layer perceptron for learning the attention scores of node tokens of the same channel for different nodes , respectively represent the node tokens of the th channel for the i-th and j-th nodes.

[0080] Secondly, sparsify the attention scores through a non-linear function:

[0081]

[0082] Among them, represents the sparse attention score, represents the softmax function, represents the set of the top k largest nodes in represents the attention score between the node token of the th channel of the i-th node in the graph data and the node tokens of the th channel of the remaining nodes.

[0083] Combine with the original adjacency matrix to form an updated adjacency matrix through residual connection:

[0084]

[0085] Among them, the finally generated is the graph structure corresponding to the set of node tokens .

[0086] S23. Combine the node token sets of all nodes in each split channel and its corresponding learned adjacency matrix together to form a learnable key graph patch, denoted as:

[0087]

[0088] Among them, represents the key graph patch of the th channel. This key graph patch is not only applicable to single-domain tasks but also can achieve migration optimization in a multi-domain environment. By combining each set of node tokens with its corresponding graph structure, the representation ability of the graph patch is enhanced. Each graph patch contains information about a certain part of node features and their context relationships in the graph structure. At this stage, the present invention particularly focuses on the migration adaptability of graph patches between domains, that is, how to make the graph patches generated in different domains perform consistently in the same latent space, significantly reducing the distribution shift between domains. S3. Encode the key graph patches of each channel in the graph data using a single-channel graph patch encoding module, and then use a graph patch aggregation module to aggregate the encoding results of the key graph patches of all channels as the graph data representation.

[0089] In this step, an optional implementation method is as follows:

[0090] S31. Based on and , encode the graph patch of the th channel:

[0091]

[0092]

[0093] Among them, is the graph patch representation of the th channel based on the original adjacency matrix, which considers the influence of the original graph structure on graph patch encoding; is the graph patch representation of the th channel based on the updated adjacency matrix; is a graph neural network;

[0094] S32. Fuse the two parts of constructed graph patches:

[0095]

[0096] Among them, is a multi-layer perceptron, represents concatenation, represents a non-linear function with a value range of (0,1), Represents the single-channel graph patch encoding. The graph patch encodings of all channels constitute .

[0097] S33, using a Transformer encoder to aggregate graph patches of different channels:

[0098]

[0099]

[0100]

[0101] Among them, is the multi-head attention mechanism; is layer normalization, used to provide training stability; is the feed-forward network; is the pooling layer, pooling in the direction of K graph patches; finally obtaining as the final graph data representation; are the intermediate results calculated by layer normalization respectively.

[0102] To handle the possible differences in channel characteristics between data in different domains, a cross-channel information aggregation mechanism is introduced through the Transformer encoder to uniformly model the representations of graph patches from different channels. This method can capture the interaction relationships between different channels, thereby improving the applicability of the final graph representation.

[0103] S4, based on the graph data representation results, introduce the feature masking restoration task and the graph neighbor context prediction task to pre-train the graph patch extraction module, the single-channel graph patch encoding module, and the graph patch aggregation module. In this step, use the pre-training method to train the learnable graph patch extraction module, the single-channel graph patch encoding module, and the graph patch aggregation module simultaneously. Specifically, use two loss functions of the feature masking restoration task and the graph neighbor context prediction task and , where the feature masking restoration task is the reconstruction loss for node feature recovery, used to ensure that the model can retain important information in the original node features; the graph neighbor context prediction task is the loss for context prediction, used to enhance the model's learning ability of local structural relationships. In the optimization process of the loss function, the present invention particularly focuses on the consistency between domains, ensuring that the model can adapt to the data in both the source domain and the target domain during training.

[0104] Feature masking and restoration task: Mask part of the node features of the graph data, generate a graph data representation based on the masked graph data; use the first decoder to restore the masked node features according to the graph data representation; calculate the feature masking and restoration loss by combining the true masking result and the restoration result:

[0105]

[0106] Among them, represents the mask of the j-th dimension feature of the i-th node, the dimension of the original node feature, represents the number of masked features, represents the number of nodes in the graph data, are the true masking result and the restoration result of the j-th dimension feature of the i-th node respectively. In this task, by masking some node features, the model is made to make predictions, and then the difference between the true value and the predicted value is measured to achieve the effect of training the model. Here, the first decoder can adopt a multi-layer perceptron.

[0107] Graph neighbor context prediction task: Generate a graph data representation based on the original graph data; use the second decoder to predict the mean feature vector of the p-th order neighbor nodes of each node in the graph data, where p = 1, 2,..., P; P represents the set maximum order; calculate the graph neighbor context prediction loss by combining the true result and the predicted result:

[0108]

[0109] Among them, represent the true result and the predicted result of the mean feature vector of the p-th order neighbor nodes of node i respectively. Here, the second decoder can adopt the same structure as the first decoder, such as a multi-layer perceptron, and the two decoders are independent of each other.

[0110] Use the weighted result of the two losses as the final loss to update the trainable parameters of the model. Each module after pre-training is combined as a cross-domain graph model for extracting the target domain graph data representation in practical applications to achieve the downstream tasks of the target domain graph data. The present invention can effectively improve the migration performance of the graph model and significantly reduce the learning obstacles caused by domain differences.

[0111] The target domain referred to in the present invention can be the same as or different from the source domain. For example, the target domain can be molecular graph data, and the downstream task of the target domain graph data is a molecular graph classification task, which classifies the molecular graphs into corresponding categories according to chemical properties. In order to obtain better representation results, the pre-trained model can also be fine-tuned using the target domain graph data on the corresponding downstream tasks.

[0112] To verify the effectiveness of the present invention, in this embodiment, graph datasets from multiple source domains are used for pre-training, and the proposed method is tested on molecular graph datasets and citation graph network datasets to demonstrate the generality of the method. Among them, the molecular graph dataset contains a dataset of 2 million molecules extracted from the ZINC commercial compound database as the pre-training dataset; the molecular graph dataset also contains Tox21, Toxcast, Sider, ClinTox, MUV, HIV, BBBP, BACE, etc. as downstream datasets. The citation graph network dataset contains Arxiv and DBLP, with 169,343 and 28,702 articles respectively, where DBLP is used as the downstream test dataset. The social network dataset uses the Flickr dataset, which contains 105,938 user nodes.

[0113] Evaluation metrics: The evaluation metric method of the present invention is to evaluate according to the effect of the pre-trained model in the downstream. Compared with the method without pre-training, the higher the effect of pre-training, the better the performance. Since the present invention aims to construct a cross-domain model, not only upstream and downstream data in the same domain but also upstream and downstream data in different domains need to be used for verification. For graph classification problems, the commonly used ROC-AUC metric is used.

[0114] Comparison method: In the experiment, the existing common model GIN is selected as the basic model for comparison, and different degrees of migration are compared. The test results are shown in Table 1.

[0115] Table 1 Migration effect

[0116]

[0117] As shown in Table 1, the method proposed in the present invention demonstrates excellent capabilities in handling feature heterogeneity. Whether it is during the transfer between pre-training datasets or between pre-training and downstream datasets, it can significantly improve the model performance. This heterogeneity is mainly reflected in the distribution differences of multi-domain data and the diversity of task objectives. The method of the present invention effectively narrows the performance gap caused by these distribution differences by combining pre-training and transfer learning. Without pre-training, the performance of the traditional GIN model and the method of the present invention in downstream tasks is similar, reflecting the basic capabilities of the basic graph neural network in the end-to-end training scenario. However, the method of the present invention shows a consistent and significant performance improvement after combining pre-training. For example, on the Sider dataset, after joint pre-training (ZINC and Arxiv datasets), the method of the present invention achieved an accuracy of 53.71±0.48, which is better than 52.14±0.56 of the GIN model without pre-training; in the HIV dataset, joint pre-training increased the accuracy from 56.58±2.57 without pre-training to 60.23±1.15. The same trend is also reflected in other datasets and will not be elaborated here. These results fully demonstrate that pre-training can provide strong support for transfer learning, especially in tasks with significant data distribution differences, and its effect is more obvious.

[0118] It can also be observed from Table 1 that joint pre-training is more effective than pre-training on a single dataset. For example, on the BACE dataset, the accuracy achieved by the pre-training method based only on the ZINC dataset is 59.74±0.54, while joint pre-training combining ZINC and Arxiv further increased the accuracy to 59.75±1.02. This indicates that pre-training data from different sources can play a complementary role in the information integration process, significantly enhancing the generalization ability of the model. In addition, for some domain data (such as the DBLP dataset), even if there are significant differences in the backgrounds between pre-training datasets, the joint pre-training strategy can still achieve excellent performance, reaching an accuracy of 79.12±1.93, which is better than 74.86±2.26 of the model without pre-training. This performance further proves the applicability and advantages of the method of the present invention in cross-domain transfer tasks.

[0119] It should be noted that while achieving the above performance improvement, the method of the present invention also exhibits good robustness. For example, in some extreme scenarios, when there is a significant inconsistency in the feature distributions between the pre-training data and the downstream data, the present invention can still provide consistent performance improvement. This indicates that through the optimized design of the feature partitioning and structure learning of graph data, the method of the present invention can efficiently capture the domain-agnostic characteristics between data, thereby achieving effective transfer in multi-task scenarios. The test results in Table 1 comprehensively demonstrate the wide applicability and superior performance of the cross-domain graph model proposed by the present invention in multi-task and multi-domain scenarios. Through the joint pre-training and optimized design of data from different domains, the method of the present invention significantly improves the transfer ability of the model, laying a solid foundation for the realization of a general graph model and providing an important reference for future research on the graph model infrastructure.

[0120] Based on the same inventive concept, in this embodiment, a pre-training system for a cross-domain graph model based on learnable graph patches is further provided. The system includes:

[0121] A pre-training graph dataset acquisition module, which is used to acquire pre-training graph datasets from multiple source domains, including two or more of citation graph data, social network graph data, and molecular graph data. Each graph data consists of a node feature matrix and an original adjacency matrix;

[0122] A graph patch extraction module, which is used to split the node features of each graph data in the pre-training graph dataset into a set of node tokens for each channel, and respectively establish a graph structure based on the set of node tokens for each channel. The set of node tokens for each channel and its corresponding graph structure for that channel are used as the key graph patches for that channel;

[0123] A single-channel graph patch encoding module, which is used to encode the key graph patches for each channel in the graph data;

[0124] A graph patch aggregation module, which is used to aggregate the encoding results of the key graph patches for all channels and use the aggregation result as the graph data representation;

[0125] A pre-training module, which is used to perform pre-training on the graph patch extraction module, the single-channel graph patch encoding module, and the graph patch aggregation module based on the graph data representation result, introducing a feature masking and restoration task and a graph neighbor context prediction task. The pre-trained module combination is used as a cross-domain graph model for extracting the target domain graph data representation to implement the downstream tasks of the target domain graph data.

[0126] For the system embodiments, since they basically correspond to the method embodiments, the relevant parts can be referred to the descriptions in the method embodiments, and the implementation methods of the remaining modules will not be elaborated here. The system embodiments described above are merely illustrative. The units described as separate components may or may not be physically separated, and the components shown as units may or may not be physical units, that is, they may be located in one place or distributed to multiple network units. Some or all of the modules can be selected according to actual needs to achieve the purpose of the solution of the present invention. A person of ordinary skill in the art can understand and implement it without creative efforts.

[0127] The embodiments of the system of the present invention can be applied to any device with data processing capabilities, and the device with data processing capabilities can be a device or apparatus such as a computer. The system embodiments can be implemented by software, or by hardware or a combination of software and hardware. Taking software implementation as an example, as a logically meaningful device, it is formed by the processor of any device with data processing capabilities reading the corresponding computer program instructions in the non-volatile memory into the memory for operation.

[0128] The above-described embodiments only represent several implementation manners of the present invention, and the descriptions are relatively specific and detailed, but should not be construed as limiting the scope of the present invention. For those of ordinary skill in the art, without departing from the concept of the present invention, several modifications and improvements can still be made, and these all belong to the protection scope of the present invention.

Claims

1. A cross-domain graph model pre-training method based on learnable graph patches, characterized in that: include: (1) Obtain pre-trained graph datasets from multiple source domains, including two or more of citation graph data, social network graph data, molecular graph data, financial network graph data, power network graph data, recommendation system graph data, and transportation network graph data. Each graph data consists of a node feature matrix and an original adjacency matrix. (2) Using the graph patch extraction module, the node features of each graph data in the pre-trained graph dataset are split into node token sets of each channel. A graph structure is established based on the node token sets of each channel. The node token sets of each channel and the graph structure of the corresponding channel are used as the key graph patches of the channel. (3) Using the single-channel image patch encoding module to encode the key image patches of each channel in the image data, and then using the image patch aggregation module to aggregate the encoding results of the key image patches of all channels as the image data representation; (4) Based on the graph data representation results, the feature masking restoration task and the graph neighbor context prediction task are introduced to pre-train the graph patch extraction module, the single-channel graph patch encoding module and the graph patch aggregation module. The pre-trained module combination is used as a cross-domain graph model for extracting the representation of the target domain graph data to achieve downstream tasks of the target domain graph data.

2. The cross-domain graph model pre-training method based on learnable graph patches according to claim 1, characterized in that: The calculation process of the image patch extraction module includes: (2-1) Split the node features of each graph data into sliding window-like segments, where the node features are split into node tokens. The node token is used as the The node tokens of each channel have the same dimension; (2-2) Construct the graph structure of the corresponding channel based on the node token set of each channel disassembled: ; in, represents a neural network model for learning graph structures, represents the original adjacency matrix of the graph data, Represents the first The node token set of the channel, Represents the number of nodes in the graph data, Indicated by The generated The node token set of each channel corresponds to the adjacency matrix of the graph structure, which is used to represent a graph structure; (2-3) According to the node token set of each channel and the graph structure of its corresponding channel, the key graph patch of each channel is obtained: ; in, Indicates Keymap patches for each channel.

3. The cross-domain graph model pre-training method based on learnable graph patches according to claim 2, characterized in that: The neural network model for learning graph structure is a neural network model based on the attention mechanism, and the calculation process includes: The first The attention scores between the node tokens of each channel; The attention scores are sparsely distributed through nonlinear functions, and the sparse attention scores between all nodes in the graph data constitute the first The sparse attention adjacency matrix of channels; The data of the figure The sparse attention adjacency matrix of the channels is residually connected with the original adjacency matrix to form the The node token set of each channel corresponds to the adjacency matrix of the graph structure.

4. The cross-domain graph model pre-training method based on learnable graph patches according to claim 1, characterized in that: The calculation process of the single-channel image patch encoding module includes: (3-1) The node token set of the channel and the original adjacency matrix are used as the input of the graph neural network to generate the first Channel-wise patch representation ; And, the The node token set of the channel is the same as the The adjacency matrix of the graph structure corresponding to the node token set of the channel is used as the input of the graph neural network to generate the first Channel-wise patch representation ; (3-2) The two image patch representations are fused to obtain a single-channel image patch encoding. After traversing all channels, the image patch encoding of all channels is obtained.

5. The cross-domain graph model pre-training method based on learnable graph patches according to claim 4, characterized in that: The process of fusing two image patch representations is as follows: ; in, represents a multi-layer perceptron, Indicates splicing, represents a nonlinear function with a range of (0,1), Indicates Channel-wise image patch encoding.

6. The cross-domain graph model pre-training method based on learnable graph patches according to claim 1, characterized in that: The graph patch aggregation module adopts a Transformer encoder to obtain a graph data representation by aggregating graph patch encodings of all channels.

7. The cross-domain graph model pre-training method based on learnable graph patches according to claim 1, characterized in that: The feature masking restoration task includes: Masking some node features of the graph data, and generating a graph data representation based on the masked graph data; Using the first decoder to restore the obscured node features according to the graph data representation; The feature occlusion restoration loss is calculated by combining the true occlusion result and the restoration result.

8. The cross-domain graph model pre-training method based on learnable graph patches according to claim 1, characterized in that: The graph neighbor context prediction task includes: Generate graph data representation based on original graph data; Using the second decoder to predict the feature mean vector of the p-order neighbor nodes of each node of the graph data according to the graph data representation, p=1,2,…,P; where P represents the set maximum order; The graph neighbor context prediction loss is calculated by combining the true results and the predicted results.

9. The cross-domain graph model pre-training method based on learnable graph patches according to claim 1, characterized in that: The target domain graph data is molecular graph data, and the downstream task of the target domain graph data is a molecular graph classification task.

10. A cross-domain graph model pre-training system based on learnable graph patches, used to implement the method of claim 1; characterized in that: The system comprises: A pre-trained graph dataset acquisition module, which is used to acquire pre-trained graph datasets of multiple source domains, including two or more of reference graph data, social network graph data, and molecular graph data, each of which consists of a node feature matrix and an original adjacency matrix; A graph patch extraction module is used to split the node features of each graph data in the pre-trained graph dataset into node token sets of each channel, and establish a graph structure based on the node token sets of each channel. The node token sets of each channel and the graph structure of the corresponding channel are used as the key graph patches of the channel. A single-channel image patch encoding module, which is used to encode the key image patches of each channel in the image data of the encoding module; A graph patch aggregation module, which is used to aggregate the encoding results of key graph patches of all channels and use the aggregation results as graph data representation; A pre-training module is used to pre-train the graph patch extraction module, the single-channel graph patch encoding module and the graph patch aggregation module based on the graph data representation results. The pre-trained module combination is used as a cross-domain graph model for extracting the representation of the target domain graph data to achieve downstream tasks of the target domain graph data.

Citation Information

Patent Citations

  • Traffic state completion prediction method based on multi-source data pre-training

    CN119202772A

  • Training image-to-image translation neural networks

    US20200160113A1