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:
ModuleCross-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).