lgatr.layers.slim_layers.SlimCrossAttention

class lgatr.layers.slim_layers.SlimCrossAttention(q_v_channels, kv_v_channels, q_s_channels, kv_s_channels, num_heads, attn_ratio=1, dropout_prob=None)[source]

Bases: Module

Cross-attention for Lorentz vectors and scalar features.

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

  • kv_v_channels (int) – Number of key/value vector channels.

  • q_s_channels (int) – Number of query scalar channels.

  • kv_s_channels (int) – Number of key/value 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_q, vectors_kv, scalars_q, scalars_kv, **attn_kwargs)[source]

Apply cross-attention.

Parameters:
  • vectors_q (Tensor) – Query Lorentz vectors of shape (..., items_q, 4, q_v_channels).

  • vectors_kv (Tensor) – Key/value Lorentz vectors of shape (..., items_kv, 4, kv_v_channels).

  • scalars_q (Tensor) – Query scalar features of shape (..., items_q, q_s_channels).

  • scalars_kv (Tensor) – Key/value scalar features of shape (..., items_kv, kv_s_channels).

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

Return type:

tuple[Tensor, Tensor]

Returns:

  • outputs_v – Lorentz vectors of shape (..., items_q, 4, q_v_channels).

  • outputs_s – Scalar features of shape (..., items_q, q_s_channels).