lgatr.layers.attention.cross_attention.CrossAttention
- class lgatr.layers.attention.cross_attention.CrossAttention(config, primitives)[source]
Bases:
ModuleL-GATr cross-attention.
Constructs queries, keys, and values, computes geometric attention, and projects linearly to outputs.
- Parameters:
config (
CrossAttentionConfig) – Attention configuration.primitives (
PrimitivesConfig) – LGATr primitives configuration.
- forward(multivectors_q, multivectors_kv, scalars_q=None, scalars_kv=None, **attn_kwargs)[source]
Compute cross-attention.
Queries come from
multivectors_qand keys/values frommultivectors_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.