ML Attention - Heungwoo/research GitHub Wiki

ML Foundations — Attention Variants

Part of the ML Foundations section. Compiled April 2026. Scope: every attention flavor that shows up inside a modern VLA or Qwen-class VLM, with diagrams and concrete VLA/Qwen3/Qwen3.5 cross-references.


1. Why attention has a zoo of variants

A standard transformer block spends most of its compute in attention. That compute scales as O(N²·d) with sequence length N. Over the 2020–2026 period, almost every architectural tweak has been motivated by one of three pressures:

  1. KV-cache memory grows linearly with layers × heads × seq-len. Inference servers hit this wall before they hit FLOPs.
  2. Quadratic cost is fatal for long-context (100k+ tokens) video/document workloads. Softmax attention does not get cheaper with sparsity unless you rewrite the kernel.
  3. Training stability — when heads get wide (head_dim ≥ 128), logits blow up, gradients spike, and training diverges at large scale. This is what QK-Norm fixes.

Every variant below is a point in this trade-off space.

flowchart TB
  ROOT["Attention"]
  ROOT --> SOFT["Softmax family<br/>quadratic O N squared"]
  ROOT --> LIN["Linear / kernel family<br/>linear O of N"]

  SOFT --> SG1["Head grouping<br/>MHA, MQA, GQA"]
  SOFT --> SG2["Masking<br/>Causal, Bidirectional,<br/>Sliding-window, Block-causal"]
  SOFT --> SG3["Cross-attention"]
  SOFT --> SG4["Stability<br/>QK-Norm"]
  SOFT --> SG5["Kernel<br/>FlashAttention-2 / 3"]

  LIN --> LG1["Linear attention<br/>Performer, Linformer"]
  LIN --> LG2["State-space<br/>Mamba, Mamba-2"]
  LIN --> LG3["Delta rule<br/>DeltaNet,<br/>Gated DeltaNet"]
  LIN --> LG4["Gated softmax<br/>Gated Attention"]

  SG1:::shipped
  SG3:::shipped
  SG4:::shipped
  LG3:::shipped
  LG4:::shipped
  classDef shipped fill:#fff3b0,stroke:#b58900,color:#333
Loading

Legend. Yellow-filled boxes are actually shipped in a 2026 VLA or Qwen release (details in §5): GQA is used in Qwen3 dense and MoE; cross-attention connects GR00T's DiT to its VLM; QK-Norm was added in Qwen3 for training stability; Gated DeltaNet and Gated Attention together form the Qwen3-Next / Qwen3.5 hybrid.


2. The math baseline — full softmax attention

For reference. Given token matrix X ∈ R^{N×d}, project to query/key/value:

Q = X W_Q       K = X W_K       V = X W_V       (each ∈ R^{N×d})
A = softmax(Q Kᵀ / √d_head)                      ∈ R^{N×N}
Y = A V                                          ∈ R^{N×d}

Diagram:

flowchart LR
  X["Input tokens<br/>X in R^(N x d)"] --> QP["W_Q"] --> Q["Q"]
  X --> KP["W_K"] --> K["K"]
  X --> VP["W_V"] --> V["V"]
  Q --> DOT["Q · K_transposed<br/>divided by sqrt d_head"]
  K --> DOT
  DOT --> MASK["optional mask<br/>causal, local, or cross"]
  MASK --> SM["softmax<br/>row-wise"] --> A["attention matrix<br/>shape N by N"]
  A --> MUL["attention · V"]
  V --> MUL
  MUL --> Y["Output Y"]
Loading

Everything below is a modification of this pipeline.


3. Softmax family — the variants

3.1 MHA / MQA / GQA — what differs is the KV head count

These three are algebraically the same, and differ only in how many KV heads exist relative to Q heads. The one driver is KV-cache memory; the question is how much quality you trade for how much cache savings.

flowchart TB
  subgraph MHA["MHA Multi-Head vanilla"]
    Q1["Q heads = 16"] --> AT1["attention"] --> O1["out"]
    K1["K heads = 16"] --> AT1
    V1["V heads = 16"] --> AT1
  end
  subgraph GQA["GQA Grouped-Query used in Qwen3 and Llama 3"]
    Q2["Q heads = 16"] --> AT2["attention"] --> O2["out"]
    K2["K heads = 8<br/>each shared by 2 Q heads"] --> AT2
    V2["V heads = 8"] --> AT2
  end
  subgraph MQA["MQA Multi-Query used in PaLM and Falcon"]
    Q3["Q heads = 16"] --> AT3["attention"] --> O3["out"]
    K3["K heads = 1<br/>shared by ALL Q heads"] --> AT3
    V3["V heads = 1"] --> AT3
  end
Loading

Intuition. A KV head is a look-up table of what-to-retrieve. Full MHA gives every query head its own table (expensive). MQA gives all query heads one shared table (cheap, small quality drop). GQA is the middle: Q heads are grouped, and each group shares one KV head.

Savings formula: KV cache size scales with n_kv_heads × head_dim × seq_len × n_layers. Going from MHA-16 → GQA-8 halves cache; GQA-8 → MQA-1 is another 8×.

Model Q heads KV heads Ratio Head_dim
GPT-2 12 12 1:1 (MHA) 64
Llama-2-7B 32 32 1:1 (MHA) 128
Llama-3-8B 32 8 4:1 (GQA) 128
PaLM 32 1 32:1 (MQA) 128
Qwen3-0.6B 16 8 2:1 (GQA) 128
Qwen3-8B 32 8 4:1 (GQA) 128
Qwen3-32B 64 8 8:1 (GQA) 128
Qwen3-235B-A22B MoE 64 4 16:1 (GQA) 128

Source: HuggingFace config.json files for each model.

3.2 Causal vs bidirectional masking

The mask is applied to Q·Kᵀ before softmax. The attention code itself is unchanged.

flowchart LR
  subgraph C["Causal: decoder / LLM"]
    direction TB
    c1["token i sees tokens 0..i only<br/>triangular mask"]
  end
  subgraph B["Bidirectional: encoder / DDVLA actions"]
    direction TB
    b1["token i sees ALL tokens<br/>no mask"]
  end
  subgraph BC["Block-causal: π0.7"]
    direction TB
    bc1["block A -- bidirectional inside<br/>block B -- attends A and itself, bidir within B<br/>block C -- causal over A, B, C"]
  end
Loading
  • Causal is what GPT/Llama/Qwen do during generation.
  • Bidirectional is what BERT does, what DDVLA does over its action tokens ("every action token sees every other during refinement"), and what cross-attention always effectively is on the KV side.
  • Block-causal is the π0.7 pattern: observations and subgoals are bidirectional among themselves; text is causal after them; action tokens are bidirectional again. This is what enables prefix-KV caching — the non-action prefix is deterministic for a given observation, so it can be cached once and reused across every denoising step.

3.3 Sliding-window attention

Each token only attends to the K tokens before it (e.g. K=4096). Saves cost for long contexts at the price of losing distant dependencies.

flowchart LR
  T0[t0] --- T1[t1] --- T2[t2] --- T3[t3] --- T4[t4] --- T5[t5]
  T5 -. attends .-> T2
  T5 -. attends .-> T3
  T5 -. attends .-> T4
  T5 -. attends .-> T5
Loading

Mistral-7B uses this. Qwen3 does not (sliding_window: null in all configs) — they rely on YaRN + Dual Chunk Attention at inference instead.

3.4 Cross-attention (encoder–decoder)

Queries come from one stream, keys & values from another. This is the canonical way to inject a fixed-context signal (here, VLM features) into a generator (here, a DiT).

flowchart LR
  subgraph Decoder["Decoder stream: action DiT"]
    X1["state and action tokens"] --> WQ["W_Q"] --> Qd["Q"]
  end
  subgraph Encoder["Encoder stream: VLM features"]
    X2["VLM hidden states"] --> WK["W_K"] --> Kd["K"]
    X2 --> WV["W_V"] --> Vd["V"]
  end
  Qd --> ATTN["softmax Q · K_T scaled by sqrt d<br/>times V"]
  Kd --> ATTN
  Vd --> ATTN
  ATTN --> OUT["action output"]
Loading

GR00T (all versions) uses cross-attention from its DiT into the frozen/tuned VLM hidden states. RDT-1B, DexVLA, Fast-in-Slow all use some cross-attention variant. π-series does not — it uses same-stack MoE with prefix-KV instead (see Review-VLM-Action-Connection).

3.5 QK-Norm — the stability fix

The problem. When head_dim grows (128+) and training scales past ~10B parameters, the raw dot-product Q·Kᵀ produces extreme-magnitude logits. Softmax saturates, gradients vanish on inactive entries and spike on active ones, and loss diverges.

The fix. Normalize Q and K separately, per-head, before the dot-product. Qwen3 uses RMSNorm; the original QK-Norm paper (Henry et al. 2020, arXiv:2010.04245) used L2 normalization along the head dim with a learnable scale; other implementations use LayerNorm.

flowchart LR
  X["Input"] --> WQ["W_Q"] --> Qraw["Q_raw"] --> QN["RMSNorm<br/>per-head<br/>over head_dim"] --> Q["Q"]
  X --> WK["W_K"] --> Kraw["K_raw"] --> KN["RMSNorm<br/>per-head<br/>over head_dim"] --> K["K"]
  X --> WV["W_V"] --> V["V"]
  Q --> DOT["Q · K_transposed<br/>scaled by sqrt d_head"]
  K --> DOT
  DOT --> SM["softmax"] --> A["attention"]
  A --> MUL["multiply by V"]
  V --> MUL
  MUL --> OUT["Y"]
Loading

Exact code from transformers/models/qwen3/modeling_qwen3.py:

self.q_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.k_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps)
# V is NOT normalized

This is one of the two biggest architectural deltas from Qwen2.5 → Qwen3 (the other: QKV-bias removed).

3.6 FlashAttention — a kernel, not a variant

FlashAttention-2/3 is a tiled GPU kernel for the softmax-attention operation. It produces the same output as naive softmax attention, but uses ~10× less HBM traffic by keeping tiles in SRAM and recomputing softmax on-the-fly. It does not change the math; it changes how the math runs on a GPU. Every modern VLA that is trained on H100s uses it, but none list it as an architectural choice — because it isn't one.


4. Linear / gated / state-space family — when O(N²) is the enemy

When you want 1M-token context or kilohertz-rate rollout, the softmax cost is fatal. The linear family replaces softmax(QKᵀ)V with an equivalent of (Q φ(K)ᵀ) V where φ is a feature map — this factorizes into a recurrent state S:

S_t = S_{t-1} + φ(k_t) v_tᵀ          (outer-product update, O(d²))
y_t = Q_t S_t                         (O(d²))

Total cost is O(N·d²) — linear in sequence length. The price is that each update overwrites into a fixed-size d×d state, losing the associative-memory capacity of softmax attention.

Variants that matter in 2026:

4.1 DeltaNet / Gated DeltaNet

DeltaNet uses a delta-rule update inspired by Widrow–Hoff (the delta rule for linear transformers is due to Schlag et al. 2021; arXiv:2406.06484 is the hardware-efficient parallelization of it over sequence length):

S_t = S_{t-1} (I − β_t k_t k_tᵀ) + β_t v_t k_tᵀ

The write-strength β_t is data-dependent; this lets the model replace old key-value associations (erase-then-write) rather than only adding new ones — the source of DeltaNet's stronger associative recall vs. plain linear attention. Gated DeltaNet (Qwen) adds a separate scalar decay/forget gate α_t that multiplicatively shrinks the whole state S_{t-1} each step, combining delta-rule replacement with Mamba-2-style gated decay.

4.2 Gated Attention

A softmax-attention block with per-head gating on the output:

Y = softmax(QKᵀ/√d) · V
Y_gated = σ(g(X)) ⊙ Y                  (gate is a learned function of X)

The gate gives the model a way to mute an attention head dynamically, analogous to a Mixture-of-Experts router at head granularity.

4.3 Mamba / Mamba-2 / SSM

Selective state-space models. Similar O(N) asymptote, different derivation (continuous-time SSM discretized). Not used in Qwen3 but widely explored in open-source long-context LLMs.


5. What Qwen3 / Qwen3.5 / the VLAs actually use

This is the part that is novel to this wiki.

5.1 Qwen3 (dense, 0.6B–32B, April 2025)

Aspect Value Source
Type Causal self-attention arXiv 2505.09388
Head grouping GQA, 8 KV heads across all sizes HF config.json
Head dim 128 (constant) same
RoPE base θ 1,000,000 (up from 10,000 in Qwen2) tech report
Sliding window None (sliding_window: null) HF config.json
QK-Norm Yes — RMSNorm on Q and K per-head, V not normalized modeling_qwen3.py
QKV bias Removed (was present in Qwen2.5) tech report
Native context 32K–128K depending on size; 40,960 native in config, YaRN + DCA to 128K at inference tech report

Why these choices. GQA for KV-cache. QK-Norm for large-scale training stability (this is the key Qwen2.5 → Qwen3 delta). θ=1e6 to support long context via YaRN rescaling. No sliding window because YaRN+DCA is already a long-context solution that preserves full attention.

5.2 Qwen3-Next-80B-A3B and Qwen3.5 (Feb 2026)

Qwen3.5 is a real release — the flagship Qwen3.5-397B-A17B shipped Feb 16 2026, with smaller 0.8B/2B/4B/9B and 27B/35B-A3B/122B-A10B variants following. Architecturally, Qwen3.5 inherits the Qwen3-Next hybrid recipe introduced in Sep 2025.

The hybrid recipe: 3:1 Gated DeltaNet to Gated Attention.

flowchart TB
  IN["Input tokens"]
  IN --> B0["Block 1 — Gated DeltaNet then MoE"]
  B0 --> B1["Block 2 — Gated DeltaNet then MoE"]
  B1 --> B2["Block 3 — Gated DeltaNet then MoE"]
  B2 --> B3["Block 4 — Gated Attention softmax GQA then MoE"]
  B3 --> B4["Block 5 — Gated DeltaNet then MoE"]
  B4 --> DOTS["repeat the 4-block pattern 12 times"]
  DOTS --> OUT["Output"]
Loading
  • 48 total layers laid out as 12 × (3 DeltaNet → MoE blocks + 1 Gated-Attention → MoE block).
  • The softmax block uses GQA with 16 Q heads, 2 KV heads, head_dim 256, rotary dim 64.
  • MoE: 512 experts, top-10 + 1 shared, activation ratio ~1:50, 3B active of 80B.
  • Native context 262K, YaRN to ~1M.

Why hybrid. Gated DeltaNet is O(N) but has limited associative-memory capacity. Gated Attention is O(N²) but is the only operation that can do exact, content-addressable recall over an arbitrary window. The 3:1 ratio is the empirical trade-off that Qwen found.

5.3 Qwen3-VL

  • Backbone: same Qwen3 GQA + QK-Norm + RMSNorm recipe.
  • Vision tokens are concatenated into the single causal text stream — there is no cross-attention block. The LLM simply sees [image_patches || text_tokens].
  • DeepStack: multi-level ViT features are injected into multiple LLM layers, not just the first (preserves low-level detail).
  • Interleaved-MRoPE (3D RoPE): RoPE frequencies distributed across (time, height, width) interleaved, improving long-video reasoning over Qwen2.5-VL's block-partitioned MRoPE.
  • 256K native context, 1M extendable.

5.4 Attention in the major 2026 VLAs

VLA VLM-trunk attention VLM → action attention Special patterns
π0.6 / π0.7 Gemma3-4B causal self-attn; bidirectional on image tokens Same-stack MoE + prefix-KV — action expert shares stack with VLM, attends cached prefix KV Block-causal: obs & subgoal bidir; text causal after; action tokens bidir. Enables prefix-KV caching across denoising steps.
GR00T N1 Eagle-2 causal self-attn Cross-attention from DiT into layer-12 of truncated VLM select_layer=12 mid-layer tap
GR00T N1.5+ Eagle-2.5 / Cosmos-Reason / Qwen3-VL Cross-attention from DiT into last hidden layer N1.6+ adds AlternateVLDiT: every 2N-th block attends image+text; other blocks attend image only
RDT-1B Causal VLM Cross-attention Alternating Condition Injection — image and text cross-attended in alternate layers to prevent image from drowning text
Discrete Diffusion VLA (DDVLA) SigLIP + DINOv2 ViT + Llama 2 No separate connection — actions are masked tokens in the VLM's own stream Bidirectional attention over action tokens (converts causal → bidir on the action span during refinement); adaptive unmasking ordering
Fast-in-Slow LLaVA-class 32-block LLM S1 = last 2 blocks (31–32) re-run at high frequency with shared parameters Heterogeneous inputs at different frequencies (1:4 ratio, blocks=2 ablated as optimal)

5.5 AlternateVLDiT, in detail

flowchart TB
  subgraph DiT["GR00T N1.6 AlternateVLDiT, 32 blocks total"]
    direction TB
    SA["Self-attention<br/>on state and action"]
    CA["Cross-attention<br/>to VL features"]
  end
  CA -->|"block 0, 4, 8 ... every 2N blocks"| TEXT_IMG["cross-attend to<br/>TEXT and IMAGE"]
  CA -->|"block 1, 2, 3, 5, 6, 7 ..."| IMG_ONLY["cross-attend to<br/>IMAGE only"]
Loading

The mask is constructed inside the block:

image_attention_mask = image_mask & backbone_attention_mask
non_image_attention_mask = (~image_mask) & backbone_attention_mask
if idx % (2 * self.attend_text_every_n_blocks) == 0:
    # cross-attend to TEXT (+ image)
else:
    # cross-attend to IMAGE only

Reason: image patch tokens outnumber text tokens ~100:1, so giving every block equal attention over both drowns the language signal. Text is attended every 4 blocks (with N=2 default).

5.6 π-series block-causal + prefix-KV caching — why π can't use linear attention

flowchart LR
  subgraph Prefix["Prefix: computed once, cached"]
    OBS["image tokens<br/>bidirectional among themselves"] --> SG["subgoal tokens<br/>bidir within<br/>can attend obs"]
    SG --> TXT["text tokens<br/>causal after obs and subgoal"]
  end
  subgraph Action["Action tokens: re-computed every denoising step"]
    A["action tokens<br/>bidirectional among themselves<br/>attend all of Prefix"]
  end
  Prefix --> Action
Loading
  • Prefix computes its KV cache once per observation.
  • Every flow-matching denoising step re-runs only the action portion, attending into the cached prefix KV via softmax attention.
  • This requires softmax attention with an addressable cache — it does not port to linear attention, because a DeltaNet-style state is not random-access.

5.7 Bidirectional attention over action tokens in DDVLA

flowchart LR
  subgraph Causal["Standard VLM causal mask"]
    A0["a0"] --> A1["a1"] --> A2["a2"] --> A3["a3"]
    A3 -- "sees 0..3" --> A3
  end
  subgraph Bidir["DDVLA action span: mask replaced with all-ones"]
    B0["a0"] <--> B1["a1"]
    B1 <--> B2["a2"]
    B2 <--> B3["a3"]
    B0 <--> B2
    B0 <--> B3
    B1 <--> B3
  end
Loading

The discrete-diffusion refinement needs every action token to see every other, because an update to a0 should be informed by the current estimate of a3 (and vice versa). Causal masking would break this. DDVLA turns the mask off on the action span only.


6. Decision guide

If you need… Pick
Standard open-source LLM recipe, compatible with every inference server GQA + QK-Norm + pre-RMSNorm (Qwen3 / Llama 3+)
1M+ token context at affordable cost 3:1 Gated DeltaNet : Gated Attention hybrid (Qwen3-Next / Qwen3.5)
Flow-matching VLA with cacheable prefix and matched-stack action expert Block-causal softmax attention + prefix-KV (π-series)
Discrete-diffusion VLA where actions refine in parallel Bidirectional mask on action span (DDVLA)
Separate action transformer that conditions on a frozen VLM Cross-attention from DiT into VLM hidden states, with AlternateVLDiT if text tokens are drowning (GR00T N1.6+)

Links

Related pages

← Back to ML Foundations · Home

⚠️ **GitHub.com Fallback** ⚠️