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:
LightningModuleProposed 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.ReLUedge_enhance (
bool, default:True) – Applies learnable weight matrix with node-pair in output node calculation.
- 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:
LightningModuleProposed 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.ReLUnorm_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.ReLUattn_edge_enhance (
bool, default:True) – Applies learnable weight matrix with node-pair in output node calculation in MHA block.
- class graphnet.models.components.grit_layers.SANGraphHead(dim_in, dim_out, L, activation=<class 'torch.nn.modules.activation.ReLU'>, pooling)[source]¶
Bases:
LightningModuleSAN 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.ReLUpooling (
str, default:'mean') – Node-wise pooling operation. Either “mean” or “add”.