Graph learning-oriented joint task and distribution generalization method

Through the joint task and distributed generalization method for graph learning, the problem of insufficient generalization ability of graph learning models in the existing technology under the scarcity of data and distribution changes is solved, and the model is quickly adapted and efficient generalization is achieved, which is suitable for a variety of practical scenarios.

CN119962626APending Publication Date: 2025-05-09BEIJING UNIV OF POSTS & TELECOMM
View PDF 0 Cites 1 Cited by

Patent Information

Application Number
CN202510050118.3
Authority / Receiving Office
CN · China
Patent Type
Applications(China)
Current Assignee / Owner
Filing Date
2025-01-13
Publication Date
2025-05-09

AI Technical Summary

Technical Problem

The existing graph learning technology has shortcomings in task generalization and distribution generalization, especially in the case of scarcity of data and distribution changes, the generalization ability and reliability of the model are limited.

Method used

A joint task and distribution generalization method for graph learning is proposed. By obtaining the source task set, adaptive sample set and target task set corresponding to protein molecules, the training set is used to train the graph prediction model, including the input module, the refiner module and the predictor module, adaptive training and target task prediction are performed.

Benefits of technology

It realizes rapid adaptation and efficient generalization of the model in the case of scarcity of data and changes in distribution, improves the ability of task generalization and distribution generalization, and is suitable for practical scenarios such as molecular property prediction and protein function prediction.

✦ Generated by Eureka AI based on patent content.

Smart Images

  • Figure CN119962626A_ABST
    Figure CN119962626A_ABST
Patent Text Reader

Abstract

The invention discloses a graph learning-oriented joint task and distribution generalization method. The method comprises the following steps of obtaining a source task set, an adaptive sample set and a target task set corresponding to protein molecules; training the neural network model by using the training set to obtain a graph prediction model; the graph prediction model comprises an input module, a refiner module and a predictor module; using the adaptive sample set to carry out adaptability training on the graph prediction model to obtain a specific graph prediction model; and inputting the target task set into the specific graph prediction model, and outputting a protein molecule prediction result corresponding to the target task set through the specific graph prediction model. According to the method, redundant information in graph data can be reduced by extracting the task key sub-graphs, and the prediction accuracy and generalization of the model are improved.
Need to check novelty before this filing date? Find Prior Art

Description

Technical Field

[0001] The present invention relates to the technical field of graph neural network optimization, and in particular to a joint task and distribution generalization method for graph learning. Background Art

[0002] Traditional graph learning techniques have the following problems in task generalization and distribution generalization:

[0003] 1) Insufficient task generalization: Existing out-of-distribution graph learning and graph prompting methods can transfer knowledge from the source task, but when labeled samples are extremely scarce, the adaptation effect is limited, making it difficult for such methods to be applied to task scenarios with scarce data.

[0004] 2) Insufficient distribution generalization: In actual scenarios, there are often distribution differences between training data (such as drug property data) and test data (such as new drug molecules). Existing graph meta-learning methods fail to effectively handle such distribution changes, resulting in a significant drop in model performance during testing, making such methods lack reliability guarantees.

[0005] 3) Susceptible to noise interference: Current methods mostly rely on full-graph information, but actual graph data often contains a large amount of noise that is irrelevant to the task, which reduces the generalization ability of the model.

[0006] In real production environments, molecular (or protein) prediction tasks often involve extremely scarce labeled samples, and there is a significant distribution difference between labeled samples and predicted samples. All of the above problems make it difficult for existing methods to be applied to real production tasks. Summary of the invention

[0007] In view of the above-mentioned deficiencies in the prior art, the present invention provides a joint task and distribution generalization method for graph learning, which solves the problem of insufficient generalization of traditional graph analysis models in the prior art.

[0008] In order to achieve the above-mentioned invention object, the technical solution adopted by the present invention is: a joint task and distribution generalization method for graph learning, comprising the following steps:

[0009] S1, obtain the source task set, adaptation sample set and target task set corresponding to the protein molecules;

[0010] Among them, each data in the source task set is annotated with the corresponding labels of all training tasks;

[0011] S2. Use the training set to train the neural network model to obtain a graph prediction model;

[0012] The graph prediction model includes an input module, a refiner module, and a predictor module;

[0013] The input module is used to receive graph data, global prompt vectors and specific task prompt vectors corresponding to the protein sub-molecules, and process the graph data, global prompt vectors and specific task prompt vectors corresponding to the protein sub-molecules respectively to obtain graph data node embeddings corresponding to the graph data, global prompt node embeddings of the global prompt vector and specific task prompt node embeddings of the specific task prompt vector respectively;

[0014] The refiner module is used to determine a global subgraph based on graph data node embeddings and global hint node embeddings, determine a specific task hint embedding based on graph data node embeddings and specific task hint node embeddings, and determine a specific task subgraph based on the global subgraph and the specific task hint embeddings;

[0015] The predictor module is used to predict a specific task subgraph and obtain the classification result corresponding to the specific task subgraph;

[0016] S3, using the adaptation sample set to perform adaptability training on the graph prediction model to obtain a specific graph prediction model;

[0017] S4. Input the target task set into the specific graph prediction model, and output the protein molecule prediction results corresponding to the target task set through the specific graph prediction model.

[0018] The beneficial effects of the above scheme are:

[0019] (1) This paper proposes for the first time a joint task and distribution generalization framework, which solves the problem of simultaneously adapting to new tasks and handling distribution changes.

[0020] (2) The present invention significantly reduces the scale of parameters that need to be adjusted through prompt mechanism optimization, achieving the goal of quickly adapting to new tasks.

[0021] (3) The present invention reduces redundant information in graph data and improves the prediction accuracy of the model by extracting task-critical subgraphs.

[0022] (4) The present invention is applicable to practical scenarios such as molecular property prediction, protein function prediction and recommendation systems, and has broad application prospects.

[0023] Further, in S2, the input module includes a graph encoder;

[0024] The graph encoder is a GNN network.

[0025] The beneficial effect of the above further scheme is that using GNN as a graph encoder can better capture the intrinsic structure and relations in the graph data, thereby extracting more representative and discriminative features.

[0026] Further, in S2, the refiner module is an MLP network;

[0027] According to the graph data node embedding and the global hint node embedding, the global subgraph is determined, including:

[0028] S21, using the refiner module global hint node embedding for processing to obtain a first soft mask matrix;

[0029] S22, applying the first soft mask matrix to the matrix corresponding to the embedding of the graph data nodes to obtain a global subgraph;

[0030] According to the graph data node embedding and the specific task prompt node embedding, the specific task prompt embedding is determined, including:

[0031] S23, using a refiner module to process the task-specific prompt node embedding to obtain a second soft mask matrix;

[0032] S24, applying the second soft mask matrix to the matrix corresponding to the graph data node embedding to obtain the specific task prompt embedding;

[0033] Based on the global subgraph and task-specific hint embedding, the task-specific subgraph is determined, including:

[0034] S25. Apply the task-specific hint embedding to the global subgraph to obtain the task-specific subgraph.

[0035] The beneficial effect of the above further scheme is that it realizes the extraction of subgraphs and embeddings related to different tasks from the original graph data based on the MLP network, and has good feature extraction and processing efficiency.

[0036] Furthermore, in S2, the loss function L of the graph prediction model is pre for:

[0037]

[0038] Among them, T S represents the source task set, (·,·) represents mutual information, G τ ,Y τ denote the refined subgraph and label respectively, G τ =θ(G;θ τ ), G represents the original graph, γ ensures that the refined subgraph contains only the most indicative parts of the original graph, Θ(·;θ τ ) represents the refiner module, θ τ represents the refiner parameter, λ2 represents the hyperparameter, and Z pre represents the pre-training sample set, that is, the source task set, and Respectively represent graph G nFor the refined subgraph and prediction of task τ, dis(·,·) is used to measure the distance between the subgraph and the original graph as a heuristic approximation of the Kullback-Leibler (KL) divergence term, Representation graph G n The true label for task τ.

[0039] The beneficial effect of the above further scheme is that the loss function contains both the cross entropy term and the KL divergence term, which takes into account both the difference between the predicted value and the true label (through the cross entropy) and the relationship between the refined subgraph and the original graph (through the KL divergence), so that the model can learn the characteristics and structural information of the graph data more comprehensively, thereby improving the generalization ability and prediction accuracy of the model.

[0040] Furthermore, in S3, the loss function L of the specific graph prediction model is ada for:

[0041]

[0042] Among them, T T Represents the adaptation sample set.

[0043] The beneficial effect of the above further scheme is: based on the existing graph prediction model, using L ada Fine-tuning can utilize previously learned knowledge and further optimize it by adapting to the sample set, thereby achieving knowledge transfer and personalized customization of the model and reducing training time and data requirements.

[0044] Furthermore, in S2, the task-specific prompt vector is processed to obtain the task-specific prompt node embedding, and the calculation formula used is:

[0045]

[0046] Among them, H τ represents the embedding of task-specific prompt nodes, GNN(·; ·) represents the GNN network, G represents the graph data corresponding to the task-specific prompt vector, Represents the parameters of the GNN network.

[0047] Further, in S23, the refiner module is used to process the task-specific prompt node embedding to obtain a second soft mask matrix, and the calculation formula used is:

[0048]

[0049] in, represents the second soft mask matrix, MLP(·;·) represents the MLP network, Concat([·,·]) represents the concatenation operation, represents the task-specific hint node embedding of the i-th node, represents the task-specific cue node embedding of the jth node, θ mlp Represents the parameters of the MLP network;

[0050] In S24, the second soft mask matrix is ​​applied to the matrix corresponding to the graph data node embedding, and the calculation formula used is:

[0051]

[0052] in, represents the embedding of specific task prompts, ⊙ represents element-by-element matrix multiplication, and A represents the matrix corresponding to the embedding of graph data nodes;

[0053] S25, apply the task-specific hint embedding to the global subgraph to obtain the task-specific subgraph, using the calculation formula:

[0054]

[0055] Among them, G τ represents a specific task subgraph, λ1 represents the first hyperparameter, which is used to control the ratio of the global subgraph soft mask matrix and the task-specific subgraph soft mask matrix, A glo represents the global subgraph mask matrix, and X represents the node features in the graph.

[0056] Furthermore, in S2, the predictor module includes a general graph encoder and a classifier.

[0057] The beneficial effect of the above further scheme is that the predictor module includes a general graph encoder and a classifier, so that the entire module has a complete functional chain from the original graph data input to the final classification result output, can independently complete the graph data prediction task, and improves the system's integration and operability.

[0058] Furthermore, the method further comprises:

[0059] Use the Open Graph Benchmark dataset to evaluate specific graph prediction models.

[0060] The beneficial effect of the above further scheme is that due to the diversity of the data set, after the model is evaluated and trained on the data set, it can better adapt to graph data of different fields and types, improve the versatility and generalization ability of the model, and enable it to handle various complex graph structure problems in practical applications. BRIEF DESCRIPTION OF THE DRAWINGS

[0061] Figure 1 Schematic diagram of the process of a joint task and distribution generalization method for graph learning.

[0062] Figure 2 Detailed information of the Open Graph Benchmark dataset.

[0063] Figure 3 is the AUC performance of the SGP method on the SIDER and Tox21 datasets.

[0064] Figure 4 is the AUC performance of the SGP method on the MUV and ToxCast datasets.

[0065] Figure 5 Ablation results on the SIDER dataset.

[0066] Figure 6 Ablation results on the Tox21 dataset.

[0067] Figure 7 Schematic diagram of hyperparameter analysis results. DETAILED DESCRIPTION

[0068] The present invention will be further described below in conjunction with the accompanying drawings and specific embodiments.

[0069] like Figure 1 As shown in FIG. 1 , a joint task and distribution generalization method for graph learning includes the following steps:

[0070] S1, obtain the source task set, adaptation sample set and target task set corresponding to the protein molecules;

[0071] Among them, each data in the source task set is annotated with the corresponding labels of all training tasks.

[0072] For example, in protein molecule recognition and prediction, the corresponding data set is used to train the graph model. In this embodiment, task generalization refers to learning a model from a source task set and generalizing it to predict new tasks with a small number of labeled samples.

[0073] For example, a graph dataset can be divided into a training set and test set They are used for model learning and evaluation respectively, where N train and N test is the size of the training set and the test set. train ) and P(G test ) represents the graph distribution of the training set and the test set. A graph can be represented as G n =(A n ,X n ), where A n is the adjacency matrix, X n Contains node attributes. Let V n and E nRespectively represent graph G n The various classification tasks contained in the dataset can be collectively referred to as T. Each graph G in the dataset n can be labeled with all tasks, represented as In addition, data samples can be defined is a pairing of a graph and its specific task label, i.e.

[0074] Specifically, the classification task T can be divided into the source task set T for pre-training S And the adaptation sample set T T . Pre-trained sample set Contains source task set T S All graphs in and their source task labels. And, define the few-shot example set for each task τ The few-shot example set is taken from the training set G train Randomly select N ft graph, and then adapt the sample set to Finally, the test sample set Contains the test graph and its target task label. The graph model can initially be trained on the pre-training set Z pre Train on the training set Z and then adapt to the sample set Z ft Fine-tune on the test sample set Z test The joint task and distribution generalization is to consider the generalization of both tasks and distributions at the same time, where distribution generalization aims to alleviate the performance degradation caused by distribution shift in the test phase, that is, out-of-distribution generalization, that is, P(G train )≠P(G test ).

[0075] S2. Use the training set to train the neural network model to obtain a graph prediction model;

[0076] The graph prediction model includes an input module, a refiner module, and a predictor module;

[0077] The input module is used to receive graph data, global prompt vectors and specific task prompt vectors corresponding to the protein sub-molecules, and process the graph data, global prompt vectors and specific task prompt vectors corresponding to the protein sub-molecules respectively to obtain graph data node embeddings corresponding to the graph data, global prompt node embeddings of the global prompt vector and specific task prompt node embeddings of the specific task prompt vector respectively;

[0078] The refiner module is used to determine a global subgraph based on graph data node embeddings and global hint node embeddings, determine a specific task hint embedding based on graph data node embeddings and specific task hint node embeddings, and determine a specific task subgraph based on the global subgraph and the specific task hint embeddings;

[0079] The predictor module is used to predict a specific task subgraph and obtain the classification result corresponding to the specific task subgraph.

[0080] In this embodiment, the input module includes a graph encoder; the graph encoder is a GNN network.

[0081] The refiner module is an MLP network; according to the graph data node embedding and the global hint node embedding, the global subgraph is determined, including:

[0082] S21, using the refiner module global hint node embedding for processing to obtain a first soft mask matrix;

[0083] S22. Apply the first soft mask matrix to the matrix corresponding to the embedding of the graph data nodes to obtain a global subgraph.

[0084] According to the graph data node embedding and the specific task prompt node embedding, the specific task prompt embedding is determined, including:

[0085] S23, using a refiner module to process the task-specific prompt node embedding to obtain a second soft mask matrix;

[0086] S24. Apply the second soft mask matrix to the matrix corresponding to the graph data node embedding to obtain the specific task prompt embedding.

[0087] Based on the global subgraph and task-specific hint embedding, the task-specific subgraph is determined, including:

[0088] S25. Apply the task-specific hint embedding to the global subgraph to obtain the task-specific subgraph.

[0089] In this embodiment, the loss function L of the graph prediction model is pre for:

[0090]

[0091] Among them, T S represents the source task set, I(·,·) represents mutual information, G τ ,Y τ denote the refined subgraph and label respectively, G τ =θ(G;θ τ ), G represents the original graph, γ ensures that the refined subgraph contains only the most indicative parts of the original graph, Θ(·;θ τ ) represents the refiner module, θ τ represents the refiner parameter, λ2 represents the second hyperparameter, and Z pre represents the pre-training sample set, that is, the source task set, and Respectively represent graph G n For the refined subgraph and prediction of task τ, dis(·,·) is used to measure the distance between the subgraph and the original graph as a heuristic approximation of the Kullback-Leibler (KL) divergence term, Representation graph G n The true label for task τ.

[0092] In this embodiment, the specific task prompt vector is processed to obtain the specific task prompt node embedding, and the calculation formula used is:

[0093]

[0094] Among them, H τ represents the embedding of task-specific prompt nodes, GNN(·; ·) represents the GNN network, G represents the graph data corresponding to the task-specific prompt vector, Represents the parameters of the GNN network.

[0095] In this embodiment, in S23, a refiner module is used to process the embedding of the specific task prompt node to obtain a second soft mask matrix, and the calculation formula used is:

[0096]

[0097] in, represents the second soft mask matrix, MLP(·;·) represents the MLP network, Concat([·,·]) represents the concatenation operation, represents the task-specific hint node embedding of the i-th node, represents the task-specific cue node embedding of the jth node, θ mlp Represents the parameters of the MLP network;

[0098] In S24, the second soft mask matrix is ​​applied to the matrix corresponding to the graph data node embedding, and the calculation formula used is:

[0099]

[0100] in, represents the embedding of specific task prompts, ⊙ represents element-by-element matrix multiplication, and A represents the matrix corresponding to the embedding of graph data nodes;

[0101] S25, apply the task-specific hint embedding to the global subgraph to obtain the task-specific subgraph, using the calculation formula:

[0102]

[0103] Among them, G τrepresents a specific task subgraph, λ1 represents the first hyperparameter, which is used to control the ratio of the global subgraph mask matrix and the task-specific subgraph mask matrix, A glo represents the global subgraph mask matrix, and X represents the node features in the graph.

[0104] In this embodiment, the predictor module includes a general graph encoder and a classifier.

[0105] Exemplarily, the refiner module can be represented as θ(·; θ τ ):G=(A,X)→G τ =(A τ ,X), where θ τ is a task-specific refiner parameter. The ideal subgraph should be able to maximize the retention of information related to the corresponding task while minimizing redundant information. Therefore, the information bottleneck theory can be used to guide the training of the refiner. Therefore, the optimization objective can be expressed as: Make I(G τ ,G)≤γ,G τ =θ(G;θ τ ). Therefore, the variational approximation method can be used to transform the expression formula of the optimization objective into the loss function L of the graph prediction model. pre The predictor module can include a task-general graph encoder and a task-specific classifier to generate the final model prediction. Formally, once the task-specific subgraph G is obtained τ , which will be fed into a graph neural network encoder GNN(·;φ) shared by all tasks and a pooling layer to extract task-specific graph embeddings where φ is the parameter of the graph encoder. Finally, the prediction of the task is achieved by embedding the graph into Input to a two-layer task-specific multilayer perceptron MLP (·; φ τ ) to obtain: It can be simplified as:

[0106]

[0107] Optionally, during the pre-training phase, the parameters θ and φ are S , thereby learning general knowledge that will be inherited during the adaptation phase. In contrast, task-specific cues and the parameters of the task-specific classifier It is dedicated to learning task-specific knowledge and will be discarded in subsequent stages.

[0108] S3. Use the adaptation sample set to perform adaptive training on the graph prediction model to obtain a specific graph prediction model.

[0109] In this embodiment, the loss function L of the specific graph prediction model is ada for:

[0110]

[0111] Among them, T T Represents the adaptation sample set.

[0112] S4. Input the target task set into the specific graph prediction model, and output the protein molecule prediction results corresponding to the target task set through the specific graph prediction model.

[0113] In this embodiment, the method further includes:

[0114] Use the Open Graph Benchmark dataset to evaluate specific graph prediction models.

[0115] For example, this example uses four widely used small sample molecular property prediction datasets from the Open Graph Benchmark dataset, which are widely used in graph learning research. Each dataset contains multiple tasks, which are used to evaluate the task generalization ability of the model. The statistical information of the Open Graph Benchmark dataset can be as follows: Figure 2 As shown, Figure 2 Detailed information of the Open Graph Benchmark dataset.

[0116] For example, the joint task and distribution generalization method (SGP method) for graph learning in this embodiment can be compared with three types of baseline methods. The comparison results are as follows: Figure 3 and Figure 4 As shown, Figure 3 is the AUC performance of the SGP method on the SIDER and Tox21 datasets, Figure 4 is the AUC performance of the SGP method on the MUV and ToxCast datasets. Figure 3 and Figure 4 It can be seen that the average AUC of the SGP method in the in-distribution (ID) and out-of-distribution (OOD) scenarios is improved by 9.94% and 3.69% respectively. This shows that the SGP method improves the task generalization ability of the model on both ID and OOD data without sacrificing either aspect. This performance improvement demonstrates the effectiveness of the subgraph hint and refinement module of the SGP method in task and distribution generalization.

[0117] On the SIDER and ToxCast datasets, the SGP method has a more significant improvement over other baseline methods. This is because these two datasets contain more tasks, providing richer global knowledge, helping the model to better capture and learn adaptability. On the other hand, the Tox21 and MUV datasets have fewer tasks, and the improvement in the generalization ability of the model is relatively small.

[0118] In terms of task generalization, meta-learning methods such as Meta-MGNN and PAR have performed well, especially on datasets with a large number of tasks. However, these methods have high computational costs. In contrast, the SGP method achieves better performance while maintaining low computational costs through a hint mechanism.

[0119] Although distribution generalization methods (such as DIR and GSAT) can alleviate the problem of out-of-distribution generalization, they do not actively guide the model to learn general knowledge of the task, and lack efficient parameter adaptation methods. These methods usually require retraining all parameters, which may lead to overfitting risks in the case of limited data. The SGP method only updates a small number of parameters through a prompt mechanism, so it has good generalization performance even in scenarios with limited data.

[0120] Therefore, it can be seen that the SGP method in this embodiment has significant advantages in task and distribution generalization while maintaining low computational and time costs.

[0121] Optionally, to verify the rationality of the SGP method design, the complete SGP method can be compared with the following abridged version to observe the impact of each module on the final performance:

[0122] V1: Use only the predictor, without the refiner in both pre-training and adaptation stages.

[0123] V2: Directly fine-tune the refiner in the adaptation stage without using the hint vector.

[0124] V3: Only task-specific cues are used during pre-training The target task-specific cues were randomly initialized during the adaptation phase.

[0125] V4: Only global hint p is used in pre-training glo ,The target task specific cues are initialized using the global cues in the adaptation phase.

[0126] V5: Replace the hint vector with a linear transformation matrix.

[0127] Ablation experiments are performed on the SIDER and Tox21 datasets, and the effects of task and distribution generalization are evaluated in 5-shot, 25-shot, and 50-shot settings, respectively.

[0128] like Figure 5 As shown, Figure 5 is the ablation result on the SIDER dataset. Figure 6 Ablation results on the Tox21 dataset.

[0129] Depend on Figure 5 and Figure 6 It can be seen that: 1. The complete model outperforms all pruned versions: The complete SGP model outperforms all pruned versions in all experimental settings, indicating that each module in the model plays an important role in performance improvement. 2. The necessity of the refiner: Compared with the V1 version without the refiner, the complete SGP model improves by an average of 13.38% on the SIDER dataset and 9.45% on the Tox21 dataset. This shows that task-specific subgraph extraction plays an important role in improving model performance. 3. The effectiveness of the refiner: Compared with V1 and V2, the refiner can improve the performance of the predictor even without using the hint strategy, with an average improvement of 0.90% and 1.64% on the SIDER and Tox21 datasets, respectively. This further proves the effectiveness of the refiner in task-specific subgraph extraction. 4. The impact of global and task-specific hints: By comparing the performance of V2, V3, and V4, it can be found that both global and task-specific hints contribute significantly to model performance. These hints not only help the model capture global knowledge in the pre-training stage, but also improve the ability to extract task-specific subgraphs in the adaptation stage. 5. Importance of hint vectors: Compared with replacing hint vectors with linear transformations in V5, the performance of the SGP model on the SIDER dataset dropped by an average of 5.45% and on the Tox21 dataset by 7.73%. This shows that hint vectors play an important role in effectively guiding the refiner and enhancing task generalization capabilities. These ablation experiment results demonstrate the rationality of the design of each component in the SGP model, especially the refiner and hint vectors play a key role in improving the performance of task and distribution generalization.

[0130] Optionally, in order to evaluate the sensitivity of the hyperparameters λ1 and λ2, experimental results with different λ1 and λ2 values ​​can be presented on the SIDER and Tox21 datasets, respectively. Figure 7 As shown, Figure 7 Figure 1 is a schematic diagram of the hyperparameter analysis results. Keep one of λ1 and λ2 fixed and change the other value in the range of {0.01, 0.1, 1, 10, 100}. Figure 7It can be seen that a value of λ1 that is too large or too small will lead to a decrease in model performance. For example, on the SIDER dataset, when λ1 decreases from 1 to 0.01, the performance in the 5-shot and 25-shot scenarios decreases by 4.50% and 2.02%, respectively; and when λ1 increases from 1 to 100, the performance decreases by 5.23% and 3.42%, respectively. This shows that the imbalance of the model's focus, whether it is over-emphasizing the global graph or the task-specific subgraph, will lead to a decrease in effect. In addition, as λ2 increases, the performance of the model gradually improves and stabilizes. This trend is particularly evident in the few-sample scenario. For example, when λ2 increases from 0.01 to 1, the performance in the 5-shot scenario of the SIDER and Tox21 datasets increases by 2.15% and 1.92%, respectively. This shows that appropriately increasing the proportion of the regularization term in the loss function can improve the generalization ability of the model.

[0131] Optionally, the SGP model in this embodiment targets the challenges of few-sample task generalization and distribution generalization in the graph field, and can also be applied to fields such as molecular property prediction (such as determination of chemical molecular properties and analysis of side effects of drug molecules).

[0132] Those skilled in the art will appreciate that the embodiments described herein are intended to help readers understand the principles of the present invention, and should be understood that the protection scope of the present invention is not limited to such specific statements and embodiments. Those skilled in the art can make various other specific variations and combinations that do not deviate from the essence of the present invention based on the technical revelations disclosed by the present invention, and these variations and combinations are still within the protection scope of the invention.

Claims

1. A joint task and distribution generalization method for graph learning, characterized in that: The method comprises: S1, obtain the source task set, adaptation sample set and target task set corresponding to the protein molecules; Each data in the source task set is labeled with corresponding labels of all training tasks; S2. Use the training set to train the neural network model to obtain a graph prediction model; The graph prediction model includes an input module, a refiner module and a predictor module; The input module is used to receive graph data, global prompt vectors and specific task prompt vectors corresponding to protein sub-molecules, and process the graph data, the global prompt vector and the specific task prompt vector corresponding to the protein sub-molecules respectively, to obtain graph data node embeddings corresponding to the graph data, global prompt node embeddings of the global prompt vector and specific task prompt node embeddings of the specific task prompt vector respectively; The refiner module is used to determine a global subgraph based on the graph data node embeddings and the global hint node embeddings, determine a specific task hint embedding based on the graph data node embeddings and the specific task hint node embeddings, and determine a specific task subgraph based on the global subgraph and the specific task hint embeddings; The predictor module is used to predict the specific task subgraph to obtain a classification result corresponding to the specific task subgraph; S3, using the adaptation sample set to perform adaptability training on the graph prediction model to obtain a specific graph prediction model; S4. Input the target task set into the specific graph prediction model, and output the protein molecule prediction result corresponding to the target task set through the specific graph prediction model.

2. The method according to claim 1, characterized in that In S2, the input module includes a graph encoder; The graph encoder is a GNN network.

3. The method according to claim 2, characterized in that In S2, the refiner module is an MLP network; Determining a global subgraph according to the graph data node embedding and the global hint node embedding specifically includes: S21, using the global hint node embedding of the refiner module to process to obtain a first soft mask matrix; S22, applying the first soft mask matrix to the matrix corresponding to the embedding of the graph data nodes to obtain the global subgraph; The determining the specific task prompt embedding according to the graph data node embedding and the specific task prompt node embedding specifically includes: S23, using the refiner module to process the specific task prompt node embedding to obtain a second soft mask matrix; S24, applying the second soft mask matrix to the matrix corresponding to the graph data node embedding to obtain the specific task prompt embedding; Determining the specific task subgraph according to the global subgraph and the specific task prompt embedding specifically includes: S25. Apply the specific task prompt embedding to the global subgraph to obtain the specific task subgraph.

4. The method according to claim 3, characterized in that In S2, the loss function l of the graph prediction model pre for: Among them, T S represents the source task set, I(·,·) represents mutual information, G τ ,Y τ denote the refined subgraph and label respectively, G τ =θ(G;θ τ ), G represents the original graph, γ ensures that the refined subgraph contains only the most indicative parts of the original graph, Θ(·;θ τ ) represents the refiner module, θ τ represents the refiner parameter, λ2 represents the second hyperparameter, and Z pre represents the pre-training sample set, that is, the source task set, and Respectively represent graph G n For the refined subgraph and prediction of task τ, dis(·,·) is used to measure the distance between the subgraph and the original graph as a heuristic approximation of the Kullback-Leibler (KL) divergence term, Representation graph G n The true label for task τ.

5. The method according to claim 4, characterized in that In S3, the loss function L of the specific graph prediction model ada for: Among them, T T Represents the adaptation sample set.

6. The method according to claim 5, characterized in that In S2, the specific task prompt vector is processed to obtain the specific task prompt node embedding, and the calculation formula used is: Among them, H τ represents the embedding of task-specific prompt nodes, GNN(·; ·) represents the GNN network, G represents the graph data corresponding to the task-specific prompt vector, Represents the parameters of the GNN network.

7. The method according to claim 6, characterized in that In S23, the refiner module is used to process the specific task prompt node embedding to obtain a second soft mask matrix, and the calculation formula used is: in, represents the second soft mask matrix, MLP(·;·) represents the MLP network, Concat([·,·]) represents the concatenation operation, represents the task-specific hint node embedding of the i-th node, represents the task-specific cue node embedding of the jth node, θ mlp Represents the parameters of the MLP network; In S24, the second soft mask matrix is ​​applied to the matrix corresponding to the embedding of the graph data node, and the calculation formula used is: in, represents the embedding of specific task prompts, ⊙ represents element-by-element matrix multiplication, and A represents the matrix corresponding to the embedding of graph data nodes; In S25, the specific task prompt embedding is applied to the global subgraph to obtain the specific task subgraph, and the calculation formula used is: Among them, G τ represents a specific task subgraph, λ1 represents the first hyperparameter, which is used to control the ratio of the global subgraph mask matrix and the task-specific subgraph mask matrix, A glo represents the global subgraph mask matrix, and X represents the node features in the graph.

8. The method according to claim 7, characterized in that In S2, the predictor module includes a general graph encoder and a classifier.

9. The method according to claim 8, characterized in that The method further comprises: The Open Graph Benchmark dataset is used to evaluate the specific graph prediction model.

Citation Information

Cited By

  • Molecular representation learning distribution external generalization method based on graph neural network

    CN120354085A