lgatr.layers.slim_layers.SlimSelfAttention

class lgatr.layers.slim_layers.SlimSelfAttention(v_channels, s_channels, num_heads, attn_ratio=1, dropout_prob=None)[source]

Bases: Module

Self-attention for Lorentz vectors and scalar features.

Parameters:
  • v_channels (int) – Number of vector channels.

  • s_channels (int) – Number of scalar channels.

  • num_heads (int) – Number of attention heads.

  • attn_ratio (int) – Expansion ratio for the attention hidden channels.

  • dropout_prob (float | None) – Dropout probability.

forward(vectors, scalars, **attn_kwargs)[source]

Apply self-attention.

Parameters:
  • vectors (Tensor) – Lorentz vectors of shape (..., items, 4, v_channels).

  • scalars (Tensor) – Scalar features of shape (..., items, s_channels).

  • **attn_kwargs – Optional keyword arguments forwarded to attention.

Return type:

tuple[Tensor, Tensor]

Returns:

  • outputs_v – Lorentz vectors of shape (..., items, 4, v_channels).

  • outputs_s – Scalar features of shape (..., items, s_channels).