Adaptive machine learning method to capture long-range dependencies in graph-structured data
By employing an adaptive message passing mechanism with a filtering function, the GNN effectively addresses the challenges of capturing long-range dependencies, enhancing prediction accuracy and mitigating issues of over-smoothing, over-squashing, and underreaching.
Patent Information
- Application Number
- PCT/EP2024/066233
- Authority / Receiving Office
- WO · WO
- Patent Type
- Applications
- Current Assignee / Owner
- Priority Date
- 2023-12-22
- Filing Date
- 2024-06-12
- Publication Date
- 2025-06-26
AI Technical Summary
Conventional graph neural networks (GNNs) face challenges in capturing long-range dependencies in graph-structured data due to issues such as over-smoothing, over-squashing, and underreaching, which limit their ability to accurately predict properties of nodes, edges, and graphs.
The proposed method involves training a GNN with an adaptive message passing (AMP) mechanism that adjusts the number of message passing layers and employs a message passing filtering function to filter update messages during message passing. This allows the GNN to be sensitive to long-range dependencies while mitigating the problems of over-smoothing, over-squashing, and underreaching.
The AMP method enables the GNN to capture long-range dependencies more effectively, improving the accuracy of node, edge, and graph-level predictions while reducing computational complexity and preventing information loss.
Smart Images

Figure EP2024066233_26062025_PF_FP_ABST
Abstract
Description
Adaptive machine learning method to capture long-range dependencies in graph-structured data[oooi] The present application claims benefit of the European Patent Application no. 23219983.6, filed on December 22, 2023, which is expressly incorporated herein in its entirety by reference.Technical field
[0002] The present disclosure is directed to artificial intelligence (Al) and machine learning (ML) technologies concerned with the training and application of graph neural networks (GNN) that capture long-range dependencies in graph- structured data. More specifically, the present disclosure relates to employing such ML systems and computational intelligence for predicting certain properties of a graph or an element thereof, such as nodes or edges. Application scenarios may include medical and pharmaceutical applications of molecules, as well as healthcare in general.Background
[0003] Neural networks used in Al systems are in principle known from the prior art. They usually comprise an input layer, which receives a dataset as input, one or more hidden layers, sometimes comprising convolutional layers, in which calculations are performed based on the input data, and an output layer which outputs a certain output .
[0004] A general difficulty lies in applying such known neural networks to graphs of data. For example, graphs may comprise a number of nodes which are interconnected by a number of edges, wherein both the nodes and the edges may be associated with some information. An example for graphs can be social networks, where the nodes can be users of the network and the edges can be connections to other users. Another example for graphs can be molecular structures, where individual atoms maybe the nodes and the bonds to other atoms of the molecules maybe the edges.
[0005] Such graphs may thus contain information about several properties of a system, for example, common interests of users in the social network, or pharmaceutical properties of a molecule. Thus, it is desirable to employ neural networks for the analysis of certain graphs.
[0006] However, conventional neural networks require input data of a predefined dimension, for examples images with a predefined size, vectors with a predefined number of elements, etc. Graphs however often have very different dimensions which cannot be amended as easily as for example cutting an image to fit a predefined input dimension.
[0007] Thus, to allow the processing of graphs in neural networks, so-called graph neural networks, GNNs, are often utilized. Such GNNs also comprise several convolutional layers, which generally utilize so-called message passing between nodes to process graph information. In other words, each layer receives a graph as input, performs message passing, meaning that messages including information regarding the nodes are exchanged between neighboring nodes. The output of each layer is then a modified graph comprising updated information for each node. This updated information is generated based on the original node information and the received information from neighboring nodes.
[0008] The updated nodes are sometimes denoted as “node embeddings”. In each layer, the nodes receiving messages from neighboring node therefore receive new information, such that after a certain number of layers (the particular number depending on the individual graph structure), each node embedding comprises some information about all nodes of the graph.
[0009] These node embeddings can then be used for predictions, such as node level predictions, i.e., estimation of a value or property of a previously unknown node, edge level predictions, i.e., estimation of a property of a certain edge, or graph level predictions, i.e., estimations regarding the properties of the entire graph.
[0010] Referring to the above-discussed examples, node-level predictions can for example be predictions about a certain user of a social network, edge-level predictionscan be predictions about properties of a certain inter-atomic bond in a molecule, and graph-level predictions can be predictions about a certain molecule, e.g., the suitability of an exemplary molecule for a certain pharmaceutical application. For example a GNN can be trained to assess the suitability of new molecules for treatment of a certain type of disease based on a training set of known molecules that have such an effect to varying degree.Summary of the invention[oon] GNNs know in the art are found to have several disadvantages.
[0012] For example, after a certain number of convolutions, the node embeddings tend to converge to one value. That is, the nodes become indistinguishable, such that graph-specific information can get lost, if too many layers are applied. This is particularly problematic when it comes to long-range dependencies: Many layers are required in order to capture any long-range dependencies between nodes that are particularly far apart in the graph, but at the same time, too many layers lead to convergence of the node embeddings. The issue of convergence is sometimes referred to as “over-smoothing” in the art.
[0013] Moreover, if many layers are utilized, as is necessary to capture long- range dependencies, the more messages need to be transmitted and “squashed” into a fixed-size node embedding. Thus, the number of messages transmitted grows exponentially with each layer of the GNN. This issue is sometimes referred to as “oversquashing” in the art.
[0014] Furthermore, if not enough message passing layers are provided in the GNN, it is impossible for a message from one node to reach another node, if the nodes are especially far apart in the graph. This issue is sometimes denoted as “underreaching”.
[0015] In this context, Barbero, Federico, et al. "Locality-Aware Graph-Rewiring in GNNs." arXiv preprint arXiv:23io.01668 (2023) propose “graph rewiring”, which mitigates oversquashing and, to some extent, underreaching by adding edges between far-away nodes.
[0016] Finkelshtein, Ben, et al. "Cooperative Graph Neural Networks." arXiv preprint arXiv:23io.01267 (2023) considers an architecture of fixed depth where nodes can adaptively decide to be listeners and / or broadcasters of messages or to isolate themselves.
[0017] However, both approaches are limited by a discrete choice of the action, such that a resulting loss function is not differentiable in these parameters. Thus, a training method based on these parameters can only be performed if a differentiable value is assumed, such that the training methods are somewhat unprecise and often fail to output an ideally trained GNN.
[0018] Faber, Lukas, and Roger Wattenhofer. "Asynchronous Neural Networks for Learning in Graphs." arXiv preprint arXiv: 2205.12245 (2022) proposes an asynchronous message passing protocol which represents an extreme version of message passing where nodes send messages in an asynchronous fashion. This can, in principle, solve oversquashing, oversmoothing, and underreaching; however, the computation becomes very expensive and a fixed amount of messages to be sent needs to be chosen by the user, similarly to the choice of the number of layers.
[0019] Liu, Juncheng, et al. "Eignn: Efficient infinite-depth graph neural networks." Advances in Neural Information Processing Systems 34 (2021): 18762- 18773 considers implicit neural networks for graphs that correspond to infinite-depth models for capturing long-range dependencies. This model simulates synchronous message passing with an infinite number of message-propagation steps, which does not address the problem of oversquashing and does not differentiate the importance of the messages.
[0020] Spinelli, Indro, Simone Scardapane, and Aurelio Uncini. "Adaptive propagation graph convolutional network." IEEE Transactions on Neural Networks and Learning Systems 32.10 (2020): 4755-4760 proposes an approach close where each node decides whether to stop propagating messages. Also in this case, a maximum number of layers / message-propagation steps must be fixed.
[0021] It is therefore an object of the present disclosure to address such issues of prior art technologies and to provide methods and systems that allow obtaining a trained GNN which is able to capture long-range dependencies in graph structures while mitigating the problems of over-smoothing, over-squashing.
[0022] In the following it is assumed that a reader skilled in the art has fundamental knowledge of Al, ML and in particular graph-based ML, autoencoders etc. Thus, for conciseness, terminology, and concepts such as neutral network types, and associated training algorithms that have been presented in relevant textbooks and review articles known to the skilled person are not defined and / or explained in detail herein. For example, it is assumed that the skilled reader knows structure and training of GNNs that comprise multiple message passing layers.
[0023] A first aspect of the disclosure relates to a computer-implemented method for training a graph neural network (GNN), the method comprising: obtaining one or more training graphs and an initial message passing filtering function for filtering update messages during message passing, selecting an initial number of message passing layers of the GNN and an initial message passing filtering function for filtering update messages during message passing, and training the GNN based on the one or more training graphs, comprising: adjusting the number of message passing layers, and adapting the message passing filtering function, e.g., based on the adjusted number of message passing layers.
[0024] As discussed in more detail below with reference to the drawings and potential application scenarios / use cases, such a method thus allows for improved training and thus an improved GNN which allows for more accurate predictions. Particularly, by means of the described method, a GNN can be obtained which may be sensitive to long-range dependencies, i.e., sensitive to interactions between far-away nodes in a graph, where the problems of oversmoothing and oversquashing can be mitigated by using a message passing filtering function to filter the messages exchanged between nodes during message passing. At the same time, the depth of the GNN can be adapted by adjusting the number of message passing layers. This allows for a mitigation of the problem of underreaching. Hence, the training method provides fora GNN which takes long-range dependencies into account while mitigating the problems of oversmoothing, oversquashing and underreaching.
[0025] For example, as explained in more detail below, aspects of the present disclosure allow, during training of the GNN, to render a loss function used for training to depend, e.g., in a differentiable manner, on the number of message passing layers and on a functional form of the message passing filtering function. This allows to use training algorithms known in the art, such as backpropagation, to optimize, in an application specific manner (e.g., specific to graph-prediction of medical properties of molecules), the number of message passing layers and a corresponding message passing filter function.
[0026] Another aspect of the disclosure relates to a method for node classification comprising obtaining an input graph characterizing relationships between a plurality of nodes and classifying one or more nodes of the input graph by inputting the input graph to a graph neural network, GNN, trained by a method for training the GNN disclosed herein.
[0027] Another aspect of the disclosure relates to a method for edge classification comprising obtaining an input graph characterizing relationships between a plurality of nodes and classifying one or more edges of the input graph by inputting the input graph to a graph neural network, GNN, trained by a method for training the GNN disclosed herein.
[0028] Another aspect of the disclosure relates to a method for graph classification comprising obtaining an input graph characterizing relationships between a plurality of nodes and classifying the input graph by inputting the input graph to a graph neural network, GNN, trained by a method for training the GNN disclosed herein.
[0029] In further aspects, the present disclosure also relates to computing devices and computer programs adapted to implement and / or carry out the various computer-implemented methods disclosed herein.
[0030] Aspects of the present disclosure can be used in a variety of applications including, but not limited to, several anticipated use cases in drug development, material informatics, and medical / healthcare.Brief description of the drawings
[0031] Various aspects of the present disclosure are described in more detail in the following by reference to the accompanying figures.
[0032] Fig. 1 illustrates a flow chart showing an exemplary computer- implemented method for training a GNN according to aspects of the present disclosure.
[0033] Fig. 2 illustrates a flow chart showing a method for node classification according to aspects of the present disclosure.
[0034] Fig. 3 illustrates a flow chart showing a method for edge classification according to aspects of the present disclosure.
[0035] Fig. 4 illustrates a flow chart showing a method for graph classification according to aspects of the present disclosure.
[0036] Fig. 5 shows a schematic overview of an exemplary training method for training a GNN according to aspects of the present disclosure.
[0037] Fig. 6A shows an exemplary graph usable for training a GNN according to aspects of the present disclosure.
[0038] Fig. 6B shows an exemplary message passing filtering function for the graph of Fig. 6A according to aspects of the present disclosure.
[0039] Fig. 6C shows a message passing scheme according to a conventional GNN training method.
[0040] Figs. 6D and 6E show message passing schemes of a first and second message passing layer based on the graph of Fig. 6A and the message passing filtering function of Fig. 6B according to aspects of the present disclosure.
[0041] Fig. 7A shows a computational tree diagram for a conventional GNN training method.
[0042] Fig. 7B shows a computational tree diagram for the message passing layers described with reference to Figs. 6A, 6B, 6D and 6E according to aspects of the present disclosure.
[0043] Fig. 8 shows an exemplaiy scheme of determining whether or not to change the depth of the trained GNN.
[0044] Fig. 9 illustrates an exemplary implementation of an apparatus for carrying out one or more computer-implemented methods disclosed herein.Detailed description of exemplary embodiments / implementations
[0045] In the following, some exemplaiy embodiments / implementations of the various aspects disclosed herein are described in more detail, with reference to the drawings. Naturally, the computing systems and apparatuses of the present disclosure may employ standard hardware components (e.g., a set of on-premises edge computing hardware and / or cloud-based computing resources connect to each other via conventional wired or wireless networking technology). In some implementations, application-specific hardware (e.g., circuitry for training GNMs and / or circuitry for executing a trained GNM for anomaly prediction and prevention) may also be employed. Further, such computing hardware may be configured to execute software instructions (e.g., retrieved from collocated or remote memoiy circuitry) to execute the computer-implemented methods discussed herein.
[0046] While specific feature combinations are described in the following paragraphs with respect to the exemplaiy embodiments of the present disclosure, it is to be understood that not all features of the discussed embodiments have to be presentfor realizing the disclosure, which is defined by the subject matter of the claims. The disclosed embodiments may be modified by combining certain features of one embodiment with one or more technically and functionally compatible features of other embodiments. Specifically, the skilled person will understand that features, components, processing steps and / or functional elements of one embodiment can be combined with technically compatible features, processing steps, components and / or functional elements of any other embodiment of the present disclosure as long as covered by the disclosure as specified by the appended claims.
[0047] Moreover, the various embodiments discussed herein can be implemented in hardware, software or a combination thereof. For instance, the various modules of the systems and apparatuses disclosed herein maybe implemented via application specific hardware components such as application specific integrated circuits, ASICs, and / or field programmable gate arrays, FPGAs, and / or similar components and / or application specific software modules being executed on multipurpose data and signal processing equipment such as CPUs, DSPs and / or systems on a chip, SOCs, or similar components or any combination thereof.
[0048] For instance, the various computing (sub)-systems discussed herein may be implemented, at least in part, on multi-purpose data processing equipment such as edge computing servers. Similarly, GNM training subsystems or processes discussed herein may be implemented, at least in part, on multi-purpose cloud-based data processing equipment such as a set of cloud-severs and similar technology.
[0049] Generally, neural networks, such as GNMs, are machine learning models that employ interconnected layers of nonlinear processing units to predict an output for a received input. Some neural networks include hidden layers in addition to an output layer. The output of each (hidden) layer is used as input to the next layer in the network, i.e., the next hidden layer or the output layer. Each layer of the network generates an output from a received input in accordance with current values of a respective set of network parameters (processing unit connection weights, activation function parameters, etc.). As discussed in detail in the prior art references mentioned above, some neural networks represent and process graph structures comprising nodes connected by edges. The graphs may be multigraphs in which nodes may be connectedIO by multiple edges. The nodes and edges may have associated node features and edge features. These maybe updated using node update functions and edge update functions, which may be implemented by conventional neural networks such as a multi-layer perceptron (MLP). For example, training such a graph-based neural network model e.g., via supervised learning using a training set, unsupervised learning, or combinations thereof, results - inter alia - in determining the network parameters of the neural networks implementing the node and edge update functions and / or the function of a graph encoder as accurate as possible. For example, the publication William L. Hamilton. (2020). Graph Representation Learning. Synthesis Lectures on Artificial Intelligence and Machine Learning, Vol. 14, No. 3, available at https: / / www.cs.mcgill.ca / ~wlh / grl_book / and published in Graph Representation Learning (2020) ISBN: 978-3-031-00460-5 provides a general overview of the field of graph representation learning as known to the skilled person.
[0050] As described in further detail below with reference to the drawings, as compared to the prior art technologies, the methods described herein have several advantages:Mitigating underreaching: In conventional GNNs, a common problem is called “underreaching”. This refers to the fact that to grasp the interaction between nodes that are K “hops” apart (i.e., K edges have to be passed between the nodes), it is necessary to implement at least K message passing layers. As conventional GNNs generally have a fixed number of message passing layers, this can often lead to “underreaching”, i.e., the number of message passing layers may be lower than K, such that an interaction between far-away nodes cannot be grasped by the GNN. This is for example especially problematic if large graphs, such as e.g. large molecules, shall be analyzed. Thus, by adapting the number of message passing layers according to the task at hand allows for the GNN to be trained such that the optimal number of layers is provided, such that the recognition of interaction between far-away nodes can be improved.Mitigating oversmoothing: However, another problem in conventional GNNs is the so-called oversmoothing. This refers to the fact that after a certain number of message passing layers, the node embeddings may converge to a same value, thus leading to a loss of structural information from the original graph. Thus, the further the number ofmessage passing layers is increased, the further the oversmoothing becomes an issue. This problem is mitigated by the present disclosure by (a) learning the optimal depth of the GNN and adapting the number of message passing layers to the necessary amount (see above) and (b) filtering the messages passed between nodes during message passing. By implementing a message passing filtering function, the messages passed between nodes can be filtered, such that in contrast to conventional GNNs, not all messages are passed between all neighboring nodes, but a filter is applied to select and / or weigh the messages sent by a specific node in a specific message passing layer. This ensures that the node embeddings do not converge even with an increased number of message passing layers.Mitigating oversquashing: Similarly, in conventional GNNs the problem of oversquashing arises. That is, with each message passing layer, the amount of information compressed in each node embedding increases exponentially. Consequently, as the number of message passing layers increases, an exponentially growing amount of information must be stored in each node embedding. This can lead to a substantial increase in computational complexity, and may also lead to the loss of information, because an increasing amount of information must be stored in a fixed number of bits per node. By filtering the messages passed on between neighboring nodes during message passing, the amount of information transmitted between nodes, and thus the amount of information that has to be stored in each node embedding is reduced, thus mitigating the problem of oversquashing.
[0051] Fig. 1 illustrates an exemplary computer-implemented method too for training a graph neural network (GNN) according to aspects disclosed herein. Method too may comprise obtaining (110) one or more training graphs. Method too may further comprise selecting (120) an initial number of message passing layers of the GNN and an initial message passing filtering function for filtering update messages during message passing. Method too may further comprise training (130) the GNN based on the one or more training graphs. Training (130) the GNN may comprise adjusting the number of message passing layers and adjusting the message passing filtering function, e.g., based on the adjusted number of message passing layers.
[0052] In some aspects, method 100 may further comprise selecting a family of distribution functions parametrized by one or more family parameter X and selecting an initial value Xi for the one or more family parameters X. In some aspects, adjusting the number of message passing layers may comprise computing a weighted output of the GNN based on the output of the message passing layers weighted by a distribution function of the family of distribution functions corresponding to X, determining whether to change the number of message passing layers based on the weighted output of the GNN, and changing the number of message passing layers based on the determining.
[0053] In some aspects, determining whether to change the number of message passing layers based on the weighted output of the GNN may comprise determining a threshold value associated with a predefined quantile of a distribution of the weighted output of the GNN, determining to change the number of message passing layers if the threshold value differs from a current number of message passing layers, and otherwise determining not to change the number of message passing layers.
[0054] In some aspects, changing the number of message passing layers may comprise increasing the number of message passing layers if the threshold value exceeds the current number of message passing layers, or decreasing the number of message passing layers if the current number of message passing layers exceeds the threshold value.
[0055] In some aspects, the family of distribution functions may comprise at least one of: a family of Poisson distributions, a family of Gaussian distributions, a mixture of distributions, or a combination thereof.
[0056] In some aspects, the message passing filtering function may be configured to apply a node-specific weight to update messages sent by the node to neighboring nodes.
[0057] In some aspects, the node-specific weight may comprise a real number which is greater than or equal to o and lower than or equal to 1. In some aspects, the real number may be either o or 1.
[0058] In some aspects, the node-specific weights of the message passing filtering function may depend on a current message passing layer.
[0059] In some aspects, each node may comprise an attribute vector, and each node-specific weight may comprise a real number which is greater than or equal to o and lower than or equal to 1 for each entry of the attribute vector.
[0060] In some aspects, training 130 may further comprise determining a loss function based on the weighted output of the GNN, determining a partial derivative of the loss function with respect to the one or more family parameters X, and adjusting the one or more family parameters based on the partial derivative with respect to the one or more family parameters X.
[0061] In some aspects, training 130 may further comprise determining a loss function based on the weighted output of the GNN, determining a partial derivative of the loss function with respect to one or more parameters of the message passing filtering function, and adjusting the filtering function based on the partial derivative with respect to the one or more parameters of the message passing filtering function.
[0062] Fig. 2 illustrates a flow chart showing a method 200 for node classification according to aspects of the present disclosure.
[0063] According to Fig. 2, method 200 may comprise obtaining 210 an input graph characterizing relationships between a plurality of nodes.
[0064] Method 200 may further comprise classifying 220 one or more nodes of the input graph by inputting the input graph to a graph neural network (GNN) trained by a method for training a GNN according to any aspect described herein.
[0065] Method 200 may be utilized for predicting properties of a particular node in a graph. For example, method 200 can be utilized for predicting properties of one or more particular atoms in a molecule.
[0066] Fig. 3 illustrates a flow chart showing a method 300 for edge classification according to aspects of the present disclosure.
[0067] According to Fig. 3, method 300 may comprise obtaining 310 an input graph characterizing relationships between a plurality of nodes.
[0068] Method 300 may further comprise classifying 320 one or more edges of the input graph by inputting the input graph to a graph neural network (GNN) trained by a method for training a GNN according to any aspect described herein.
[0069] Method 300 may be utilized for predicting properties of a particular edge in a graph. For example, method 300 can be utilized for predicting properties of one or more particular interatomic bonds in a molecule.
[0070] Fig. 4 illustrates a flow chart showing a method 400 for graph classification according to aspects of the present disclosure.
[0071] According to Fig. 4, method 400 may comprise obtaining 410 an input graph characterizing relationships between a plurality of nodes.
[0072] Method 400 may further comprise classifying 420 the input graph to a graph neural network, GNN, trained by a method for training a GNN according to any aspect described herein.
[0073] Method 400 may be utilized for predicting properties of an entire graph. For example, method 400 can be utilized for predicting a suitability of a molecule for a particular pharmaceutical application.
[0074] Fig. 5 illustrates a schematic overview of an exemplaiy training method 500 for training a GNN according to aspects of the present disclosure. According to Fig.5, a dataset 510 maybe provided. The dataset may comprise one or more training graphs. In some aspects, the dataset maybe labelled. For example, the dataset may comprise a number of molecules which maybe labelled with a suitability of the respective molecule for a particular pharmaceutical application.
[0075] During training 520, Adaptive Message Passing (AMP) may be performed, as shown in box 522. This AMP may comprise Depth computation 524 and Message passing with filtering 526. Based on the message passing 526, a loss function 530 may be computed. Based on the loss function 530, one or more gradients 540 may be computed. Based on these gradients, the Depth computation for the AMP may be performed. Moreover, based on the gradients, the filters for the message passing filtering function may be updated.
[0076] This training method may generate a Graph Neural Network, such as for example a Deep Graph Network 550.
[0077] Figs. 6A to 6E illustrate message filtering based on the message passing filtering function according to aspects of the present disclosure. To this end, Fig. 6A illustrates a graph 600 comprising nodes 610 denoted with the numerals 1 to 7. The numbering of nodes according to Fig. 6A is adhered to in Figs. 6B to 6E. The graph 600 further comprises edges 620, each edge connecting two neighboring nodes.
[0078] During conventional message passing, messages are generally exchanged between any two neighboring nodes, as is illustrated in Fig. 6C. Hence, for example, node 2 may send a message to node 1 and node 3 and receive a message from node 1 and node 3. Similarly, node 4 may send a message to nodes 3, 5 and 6, and receive a message from nodes 3, 5 and 6. Therefore, after a certain amount of message-passing layers, the amount of messages received by each node grows exponentially, leading to the oversquashing as discussed above. Moreover, as messages are exchanged multiple times, the node embeddings tend to converge, such that the nodes 1 to 7 would likely have the same or similar node properties after several message passing layers.
[0079] Therefore, a message passing filtering function as shown in Fig. 6B may be applied. For simplicity of description, it is assumed that each filter has either the value 1 or o. The value 1 indicates that a message is sent, and the value o indicates that no message is sent. However, it is noted that filters may also take any other values. For example, the filter may be a real number which is greater than or equal to o and lower than or equal to 1. In this case, the filter may be used as a weight for the corresponding message.
[0080] Fig. 6B illustrates the representation of an exemplary message passing filtering function according to aspects of the present disclosure. The representation of f(xi) may correspond to the message passing filtering function for the i-th node. The abscissa 1 indicates the 1-th message passing layer of the GNN. The exemplary message passing filtering function as depicted in Fig. 5 therefore applies to a graph with seven nodes (e.g., the graph depicted in Fig. 6A) and a GNN with 1=2 layers.
[0081] A message passing filtering function may take node attributes, or an intermediate node embedding as input parameter and may return, for each active layer in the network, a layer- wise and preferably differentiable mask that may be used to suppress part of or an entire message that a node might send during message passing.
[0082] As can be seen in the representation of f(xi), node 1 sends a message during the first message passing layer but does not send any message during the second message passing layer. Meanwhile, as depicted in the representation of f(x2), node 2 sends no message in the first message passing layer but does send a message in the second message passing layer. As can be seen in the representation of f(x3), node 3 does not send any message in both message passing layers.
[0083] Fig. 6D illustrates the messages passed on in the first message passing layer in the graph of Fig. 6A based on the message passing filtering function illustrated in Fig. 6B. Therefore, in the first message passing layer (1=1), messages are only passed on by nodes 1, 5, 6 and 7. Nodes 2, 3 and 4 do not pass any messages in the first message passing layer. At the same time, messages are received in the first message passing layer by nodes 2, 3, 4, 5 and 6 while nodes 1 and 7 do not receive any messages in the first message passing layer.
[0084] Fig. 6E illustrates the messages passed on in the second message passing layer in the graph of Fig. 6A based on the message passing filtering function illustrated in Fig. 6B. Therefore, in the second message passing layer (1=2), messages are only passed on by nodes 2 and 4. Nodes 1, 3, 5, 6 and 7 do not pass any messages in the second message passing layer. At the same time, messages are received in the second message passing layer by nodes 1, 3, 5 and 6 while nodes 2, 4 and 7 do not receive any messages in the second message passing layer.
[0085] As discussed above, this example is limited to the extreme case that the filters are either 1 or o. However, preferably, the filters take any real value. By using any real value, it can be ensured that the filter is differentiable, such that the gradient of the loss function can be calculated. This way, the training method can become more efficient than with discrete filters. In case of discrete filters, an estimation of the gradient maybe performed, as described in Finkelshtein et al.: “Cooperative graph neural networks”. arXiv preprint arXiv:23io.01267, 2023.
[0086] A message passed on by a node, may be weighted, i.e., multiplied by the value of the filter.
[0087] Also for simplicity, in the sample of Figs. 6A to 6E it is assumed, that each message only has one dimension. However, the node properties passed on by a message may have multiple dimensions, for example a vector or matrix. In this case, a single filter value may be applied to multiple vector or matrix elements. In other examples, the filter may have the same dimension as the node properties, such that each vector or matrix element may be filtered by a corresponding element of the filter.
[0088] During training of a GNN, the filters of the message passing filtering function may be adapted, as discussed for example with reference to Figs. 1 or 5.
[0089] Fig. 7A illustrates a computation tree diagram according to conventional message passing. In contrast, Fig. 7B illustrates a computation tree diagram as filtered according to AMP.
[0090] According to Fig. 7A, after two message passing layers, the embedding of node 3 comprises message information of nine messages in total, including messages from each node of the graph. Particularly, the embedding of node 3 comprises three representations of its own node information which was sent to neighboring nodes in the first message passing layer and is retransmitted to node 3 in the second message passing layer.
[0091] In contrast thereto, as illustrated in Fig. 7B, node 3 comprises message information of only five messages in total. Additionally, the embedding of node 3 afterthe second message passing layer does not comprise any “back-transmitted' information of itself.
[0092] Consequently, by filtering the message passing according to the abovedescribed message passing filtering function allows for a reduction in computational complexity, leading to a mitigation of oversquashing and oversmoothing.
[0093] Therefore, the node embedding for a node hvl, representing the embedding of node v in message passing layer I maybe computed as:where1and i1may be learnable functions that differ between layers I and T may be a permutation invariant function that aggregates embeddings of v’s neighbors u computed at the previous layer. aulvmay represent an edge between nodes u and v. Fi(u, I - 1) may then be the message passing filtering function defining for the t-th graph how much of each node or node embedding h^1 1is propagated through the outgoing edges in the message passing layer I and ° may represent the element- wise product.
[0094] Consequently, by filtering the messages according to the message passing filtring function, oversmoothing may be reduced, as only part of the messages are propagated, such that the embedding of nodes may not converge to the same vector even after many message passing layers.
[0095] Moreover, oversquashing may be reduced, as the amount of content passed on in each message can be reduced. Also, entire sub-trees of message-passing computation can be pruned (e.g., be shut down entirely). In this way, less information needs to be fitted into a single node embedding.
[0096] Fig. 8 shows an exemplary scheme of determining whether or not to change the depth of the trained GNN.
[0097] According to Fig. 8, an exemplary GNN 810 utilizes three message passing layers 812. According to Fig. 8, a decision should be made whether or not a fourth (or more) additional message passing layer 814 should be added to the GNN.
[0098] To this end, a family of distribution functions q(X) parametrized by one or more family parameters X may be selected, e.g. by a user. Also, an initial value Xi for the one or more family parameters X may be selected. In some aspects, the family of distribution functions may comprise a family of Poisson distributions, a family of Gaussian distributions, a mixture of distribution functions, or a combination thereof. A weighted output of the GNN based on the output of each message passing layer weighted by a distribution function of the family of distribution functions corresponding to the initial value Xi may be computed.
[0099] That is, the learnable parameters may be outputted after each message passing layer, and may be weighted, e.g., multiplied, with the probability mass function of the distribution q(X) . A weighted output of the exemplary GNN of Fig. 8 is depicted in the graph 820, showing the weighted output of layers 1 to 5 weighted with the probability mass function of q(X). However, the graph may refer to a distribution over natural numbers (e.g., all positive integers), such that the probability mass function may have a value for each natural number.
[0100] A threshold value associated with a predefined quantile of a distribution of the weighted output of the GNN may be determined. For example, the threshold value may correspond to a cutoff-value for the predefined quantile, such that, if e.g. a quantile of 0.95 is chosen, the threshold value may correspond to a value I for which 95% of the sum of weighted output values may be below the threshold.
[0101] A quantile may e.g. be predefined by a user depending on a specific task of the GNN.
[0102] It may be determined to change the number of message passing layers if the threshold value differs from a current number of message passing layers. Otherwise, it may be determined not to change the number of message passing layers.
[0103] For example, if the current number of messages passing layers is three, as in the example of Fig. 8 and the threshold value for the 0.95-quantile is four, it maybe determined to change the number of message passing layers.
[0104] In some cases, the number of message passing layers may be increased if the threshold value exceeds the current number of message passing layers. In the example of Fig. 8, since the current number of message passing layers is three and the threshold value is four, it may be determined to add an additional message passing layer.
[0105] In some examples, the number of message passing layers may be decreased if the current number of message passing layers exceeds the threshold value. For example, if the current number of messages passing layers is four and the threshold value is three, the fourth message passing layer may be removed.
[0106] The one or more family parameters X may comprise one or more learnable parameters. That is, during training, a gradient of a loss function maybe evaluated for the one or more family parameters X. If according to the gradient the loss function can be minimized by changing X, such a change may be implemented. The changed X may then lead to a change in the probability mass function of q(X) and consequently of the quantiles during the subsequent training iterations, such that the number of message passing layers, and thus the depth of the GNN, may change over the course of training the GNN.
[0107] A detailed description of deciding when to change a depth of a neural network is given in A. Nazaret and D. Blei: “Variational Inference for Infinitely Deep Neural Networks”. In: Proceedings of the 39th International Conference on Machine Learning (ICML), Baltimore, Maryland, USA, PMLR 162, 2022.
[0108] In this way, an appropriate choice of the depth of the architecture can both be learned during training of the GNN. Especially if filtering as described above is performed, the adjusting of the GNN depth allows for optimal adaptation of the GNN to the filter parameters as learned. In this way, it can be ensured that on the one hand, enough messages are exchanged to grasp long-distance dependencies in the graph (i.e.,mitigating underreaching), while at the same time avoiding oversmoothing and oversquashing. By adjusting the number of message passing layers according to the above-described methods, it can be ensured that an optimized number of layers is used which is large enough to grasp long-term dependencies, but is not too large, such that oversmoothing can also be addressed. Particularly, the filtering function and the depth of the GNN may mutually influence each other. Thus, during training, the filtering function may be adapted based on the number of message passing layers. Also, the number of message passing layers maybe adjusted based on the message passing filtering function.
[0109] Thus, combining the adjusting of GNN depth with the message passing filtering function may provide for an efficient method of training a GNN, mitigating the problems of oversmoothing, oversquashing and underreaching in an efficient manner.
[0110] Hence, in contrast to the state of the art on graph machine learning, the training methods according to the present disclosure do not require fixing a maximum depth a priori, but the depth may be learned by means of fully differentiable parameters. The adaptive message filtering mechanism (e.g., the message passing filtering function) may allow for prioritizing important messages. The combination of utilizing the message passing filtering function and the GNN depth adjusting lets the GNN model mitigate oversmoothing, oversquashing, and underreaching issues based on the task. Moreover, the training methods according to the present disclosure do not require an explicit intervention of a user to address these issues, either through novel designs or rewiring. Finally, AMP is independent of any specific message passing architecture, making it widely applicable to message-passing GraphAI technologies.[omJFig. 9 is a functional block diagram of a computing device 900 adapted for carrying out the methods disclosed herein. The computing device 900 may comprise processing circuitry 910 such as one or more processors operably connected to memory storing code comprising instructions for carrying out the methods disclosed herein, e.g., when these instructions are being executed by the processing circuitry 910. The computing device 900 may further comprise one or more communication interfaces 930 to obtain input data needed for executing the methods discussed herein (e.g., via accessing electronic data bases or by receiving sensor data). The communicationinterfaces 930 also allow to send control commands, triggers, warning messages, etc. to recipient devices for implementing some of the methods discussed above.
[0112] Thus, the computing device 900 may comprise one or more means for carrying out any method disclosed herein.Use cases
[0113] In the following, some use cases are described which may benefit from utilizing a method as described herein:
[0114] In some aspects, a method for developing a fully differentiable deep graph network that adapts its depth during training and decides which messages to send at each message-passing layer may comprise the steps of:1. Collecting a dataset comprising graphs and their respective target values to predict (i.e. a training set);2. Choosing a family of graph convolutional layers (e.g, GCN, GIN), a family of distributions <7 (A) (e.g., Poisson, Gaussian, or a mixture) and a corresponding quantile (e.g., 0.95), and the architecture of a learnable message passing filtering function f, (e.g., a multi-layer perceptron (MLP));3. Training AMP via backpropagation or a similar method to optimize the expected lower bound of the conditional likelihood function (e.g., the loss function). At each iteration: a. Adaptively selecting a depth of the network based on the quantile of the distribution <7 (A); b. Computing soft filters for the messages; c. Performing message passing to produce output predictions; d. Computing the expected lower bound to be maximizede. Computing gradients and update the parameters of the message passing architecture, the family of distributions, and the learnable message passing filtering function4. Outputting a graph machine learning model that predicts the properties of interest for a given graph or its individual entities.
[0115] Example 1: Structural dynamics of RNA molecules
[0116] In some examples, a GNN may be utilized for investigating structural dynamics of RNA molecules which might be considered for a design of a vaccine. RNA molecules are functionally rich biomolecules with increased public attention for their use in mRNA vaccines. Investigating the structural dynamics of RNA molecules is necessary for designing an effective candidate vaccine and is traditionally performed by running large-scale atomistic molecular dynamics simulations. These simulations are challenging because RNA molecules show multiple hierarchical structures due to phenomena occurring at different spatial scales (different interatomic distances). From the modeling perspective, it means that RNA’s structural dynamics are largely dependent on long-range interactions. As the performance of classical force fields is far from satisfactory, the importance of machine-learned force fields, especially those which use Deep Graph Networks (DGNs), has grown. However, as capturing non-local interactions by standard DGN architectures remains their main shortage, a method according to the present disclosure may be utilized to introduce long-range information in the training process of DGNs.
[0117] In such examples, a data set for training the GNN may comprise different molecule configurations (e.g., atom positions), and respective energy and force labels (e.g., evaluated at the desired quantum mechanical level of accuracy).
[0118] The GNN may then be trained to maximize the generalization accuracy at predicting energy and forces of a molecular system. The resulting network may then have learnt to filter out unnecessary messages as well as to determine the long-range interactions required to solve the predictive task. The GNN may also determine, attraining time, the relevant range of atomic interactions contributing to non-local effects on the RNA structure.
[0119] This may provide for a deep graph network for predicting molecular force fields which captures long-range interactions and can be used to investigate the structural dynamics of RNA and its thermodynamic properties by running atomistic molecular dynamics simulations at the desired level of accuracy.Example 2: Large Language Models (LLMs)
[0120] It is well known that current LLMs lack even basic reasoning abilities (for instance, answering to a negated sentence). Explicit modeling of the “reasoning process” of an LLM can help to improve the reasoning abilities based on our prior knowledge of how reasoning should behave. In this case, it may be desirable to teach an LLM to achieve planning abilities by computing the steps needed to solve a task. It may be assumed that a task can be solved by a proper combination of "tools”, for instance programs to be used by the planning agent. Notably, one of these tools can be an LLM agent capable of further planning for sub-tasks.
[0121] In such examples, a dataset for training may comprise a graph of tools as entities and pairwise connections if the output of a tool which can be used as input of another tool. Tools that do not accept inputs may not have incoming connections. The dataset may be made of samples of the form (input task, solution plan), where a solution plan maybe a subgraph of the original graph that describes the (ordered) computational steps to be taken to arrive at a solution for the task. The samples may be provided by the user.
[0122] The solution plan can be mathematically seen as a node classification problem: if a tool is activated at time step t, then AMP may predict 1 for this tool at message passing layer t and o otherwise. This may imply that, for this use case, the maximal depth can be provided in advance and the message filtering scheme (e.g., the message passing filtering function) may provide for the correct and interpretable behavior of the model. In particular, the filtering may produce skewed values to ensure that messages are either kept or completely shut down (from an interpretability point of view). In turns, this may ensure that, by following the path of activated messages, asubgraph of messages corresponding to the predicted solution plan may be identified after AMP is run.
[0123] This may provide for a DGN capable of producing solution plans from a given input task and may thus help an LLM agent to devise and orchestrate the solution of a complex task.Example 3: TCR-binding candidates for vaccine / drug development
[0124] One of the fundamental problems in the successful development of cancer vaccines is understanding whether a T-cell receptor (TCR), which monitors the health status of cells, identifies (or binds to) specific peptides presented on the surface of a cell. This problem is known as TCR-recognition. The use of Al-assisted tools for automated TCR-recognition relies on the ability of these systems to capture the nonlocal, long-range interactions between intra and inter-molecular systems, which can be a huge limiting step in developing Al-based TCR-recognition solutions.
[0125] Thus, a GNN may be trained based on a dataset of candidate TCRs (or parts of them) and target peptides, labelled with information about whether the pairs (TCR, peptide) bind or not and / or about their binding strength.
[0126] The GNN may then be trained to maximize the generalization accuracy at predicting binding and binding strength between a candidate TCR and a target peptide. The resulting network may learn to filter out unnecessary messages as well as to determine the long-range interactions required to solve the predictive task.
[0127] Such a GNN may be able to infer an accurate prediction whether a candidate TCR binds to a given peptide and if so, predict the magnitude of a respective binding strength.
Claims
Claims1. A computer-implemented method for training a graph neural network, GNN, the method comprising: obtaining one or more training graphs; selecting an initial number of message passing layers of the GNN and an initial message passing filtering function for filtering update messages during message passing; training the GNN based on the one or more training graphs, comprising: adjusting the number of message passing layers; and adjusting the message passing filtering function.
2. The method of claim 1, further comprising: selecting a family of distribution functions parametrized by one or more family parameters X and selecting an initial value Xi for the one or more family parameters X; and wherein adjusting the number of message passing layers comprises: computing a weighted output of the GNN based on the output of the message passing layers weighted by a distribution function of the family of distribution functions corresponding to X; determining whether to change the number of message passing layers based on the weighted output of the GNN; and changing the number of message passing layers based on the determining.
3. The method of claim 2, wherein determining whether to change the number of message passing layers based on the weighted output of the GNN comprises: determining a threshold value associated with a predefined quantile of a distribution of the weighted output of the GNN; determining to change the number of message passing layers if the threshold value differs from a current number of message passing layers; and otherwise determining not to change the number of message passing layers.
4. The method of claim 3, wherein changing the number of message passing layers comprises: increasing the number of message passing layers if the threshold value exceeds the current number of message passing layers; or decreasing the number of message passing layers if the current number of message passing layers exceeds the threshold value.
5. The method of any one of claims 2 to 4, wherein the family of distribution functions comprises at least one of: a family of Poisson distributions; a family of Gaussian distributions; a mixture of distributions; or a combination thereof.
6. The method of any one of the preceding claims, wherein the message passing filtering function is configured to apply a node-specific weight to update messages sent by the node to neighboring nodes, preferably wherein the node-specific weight comprises a real number which is greater than or equal to o and lower than or equal to 1, preferably wherein the real number is either o or 1.
7. The method of claim 6, wherein the node-specific weight of the message passing filtering function depend on a current message passing layer.
8. The method of claim 6 or 7, wherein each node comprises an attribute vector, and wherein each node-specific weight comprises a real number which is greater than or equal to o and lower than or equal to 1 for each entry of the attribute vector.
9. The method of any one of claims 2 to 8, the training further comprising: determining a loss function based on the weighted output of the GNN; determining a partial derivative of the loss function with respect to the one or more family parameters X; and adjusting the one or more family parameters based on the partial derivative with respect to the one or more family parameters X.
10. The method of any one of claims 1 to 9, the training further comprising: determining a loss function based on the weighted output of the GNN; determining a partial derivative of the loss function with respect to one or more parameters of the message passing filtering function; and adjusting the filtering function based on the partial derivative with respect to the one or more parameters of the message passing filtering function.
11. A method for node classification comprising: obtaining an input graph characterizing relationships between a plurality of nodes; and classifying one or more nodes of the input graph by inputting the input graph to a graph neural network, GNN, trained by the method of any of claims 1 to 10.
12. A method for edge classification comprising: obtaining an input graph characterizing relationships between a plurality of nodes; and classifying one or more edges of the input graph by inputting the input graph to a graph neural network, GNN, trained by the method of any of claims 1 to 10.
13. A method for graph classification comprising: obtaining an input graph characterizing relationships between a plurality of nodes; and classifying the input graph by inputting the input graph to a graph neural network, GNN, trained by the method of any of claims 1 to 10.
14. A computing device comprising at least one means for performing the method of any one of claims 1 to 13.
15. A computer program comprising one or more instructions which, when the program is executed by a computer, cause the computer to perform the method of any one of claims 1 to 13.
Citation Information
Patent Citations
EP23219983A