grit_layers

Layers for the GRIT graph-transformer architecture.

class graphnet.models.components.grit_layers.GritSparseMHA(in_dim, out_dim, num_heads, use_bias, clamp, dropout, activation=<class 'torch.nn.modules.activation.ReLU'>, edge_enhance)[source]

Bases: LightningModule

Proposed Attention Computation for GRIT.

Original code: https://github.com/LiamMa/GRIT/blob/main/grit/layer/grit_layer.py

Construct ‘GritSparseMHA’.

Parameters:
  • in_dim (int) – Dimension of the input tensor.

  • out_dim (int) – Dimension of the output tensor.

  • num_heads (int) – Number of attention heads.

  • use_bias (bool) – Apply bias the key and value linear layers.

  • clamp (float, default: 5.0) – Clamp the absolute value of the attention scores to a value.

  • dropout (float, default: 0.0) – Dropout layer probability.

  • activation (Module, default: <class 'torch.nn.modules.activation.ReLU'>) – Uninstantiated activation function. E.g. torch.nn.ReLU

  • edge_enhance (bool, default: True) – Applies learnable weight matrix with node-pair in output node calculation.

forward(data)[source]

Forward pass.

Return type:

Data

Parameters:

data (Data)

class graphnet.models.components.grit_layers.GritTransformerLayer(in_dim, out_dim, num_heads, dropout, norm=<class 'torch.nn.modules.batchnorm.BatchNorm1d'>, residual, deg_scaler, activation=<class 'torch.nn.modules.activation.ReLU'>, norm_edges, update_edges, batch_norm_momentum, batch_norm_runner, rezero, enable_edge_transform, attn_bias, attn_dropout, attn_clamp, attn_activation=<class 'torch.nn.modules.activation.ReLU'>, attn_edge_enhance)[source]

Bases: LightningModule

Proposed Transformer Layer for GRIT.

Original code: https://github.com/LiamMa/GRIT/blob/main/grit/layer/grit_layer.py

Construct ‘GritTransformerLayer’.

Parameters:
  • in_dim (int) – Dimension of the input tensor.

  • out_dim (int) – Dimension of theo output tensor.

  • num_heads (int) – Number of attention heads.

  • dropout (float, default: 0.0) – Dropout layer probability.

  • norm (Module, default: <class 'torch.nn.modules.batchnorm.BatchNorm1d'>) – Uninstantiated normalization layer. Must be either torch.nn.BatchNorm1d or torch.nn.LayerNorm.

  • residual (bool, default: True) – Apply residual connections.

  • deg_scaler (bool, default: True) – Apply degree scaling after MHA.

  • activation (Module, default: <class 'torch.nn.modules.activation.ReLU'>) – Uninstantiated activation function. E.g. torch.nn.ReLU

  • norm_edges (bool, default: True) – Apply normalization to edges.

  • update_edges (bool, default: True) – Update edges after layer.

  • batch_norm_momentum (float, default: 0.1) – Momentum of batch normalization.

  • batch_norm_runner (bool, default: True) – Track running stats of batch normalization.

  • rezero (bool, default: False) – Apply learnable scaling parameters.

  • enable_edge_transform (bool, default: True) – Apply a FC to edges at the start of the layer.

  • attn_bias (bool, default: False) – Add bias to keys and values in MHA block.

  • attn_dropout (float, default: 0.0) – Attention droput.

  • attn_clamp (float, default: 5.0) – Clamp absolute value of attention scores to a value.

  • attn_activation (Module, default: <class 'torch.nn.modules.activation.ReLU'>) – Uninstantiated activation function for MHA block. E.g. torch.nn.ReLU

  • attn_edge_enhance (bool, default: True) – Applies learnable weight matrix with node-pair in output node calculation in MHA block.

forward(data)[source]

Forward pass.

Return type:

Data

Parameters:

data (Data)

class graphnet.models.components.grit_layers.SANGraphHead(dim_in, dim_out, L, activation=<class 'torch.nn.modules.activation.ReLU'>, pooling)[source]

Bases: LightningModule

SAN prediction head for graph prediction tasks.

Original code: https://github.com/LiamMa/GRIT/blob/main/grit/head/san_graph.py

Construct SANGraphHead.

Parameters:
  • dim_in (int) – Input dimension.

  • dim_out (int, default: 1) – Output dimension.

  • L (int, default: 2) – Number of hidden layers.

  • activation (Module, default: <class 'torch.nn.modules.activation.ReLU'>) – Uninstantiated activation function. E.g. torch.nn.ReLU

  • pooling (str, default: 'mean') – Node-wise pooling operation. Either “mean” or “add”.

forward(data)[source]

Forward Pass.

Return type:

Tensor

Parameters:

data (Data)