lgatr.layers.attention.cross_attention.CrossAttention

class lgatr.layers.attention.cross_attention.CrossAttention(config, primitives)[source]

Bases: Module

L-GATr cross-attention.

Constructs queries, keys, and values, computes geometric attention, and projects linearly to outputs.

Parameters:
forward(multivectors_q, multivectors_kv, scalars_q=None, scalars_kv=None, **attn_kwargs)[source]

Compute cross-attention.

Queries come from multivectors_q and keys/values from multivectors_kv; per head, geometric attention is applied over items, the heads are concatenated, and a final linear map produces the outputs.

Parameters:
  • multivectors_q (Tensor) – Input multivectors for queries, shape (..., items_q, q_mv_channels, 16).

  • multivectors_kv (Tensor) – Input multivectors for keys and values, shape (..., items_kv, kv_mv_channels, 16).

  • scalars_q (Tensor | None) – Optional input scalars for queries, shape (..., items_q, q_s_channels).

  • scalars_kv (Tensor | None) – Optional input scalars for keys and values, shape (..., items_kv, kv_s_channels).

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

Return type:

tuple[Tensor, Tensor | None]

Returns:

  • outputs_mv – Output multivectors of shape (..., items_q, out_mv_channels, 16).

  • outputs_s – Output scalars of shape (..., items_q, out_s_channels), or None.