# Sol-Attn2

Sparse Attention with Complementary Approximation

October 2, 2026

## Abstract

We present Sol-Attn2, an 8-bit sparse attention method with constant-overhead approximation compensation that supports both training-free inference and training. Its core idea is to combine threshold-based block routing with online compression: while selecting blocks for exact attention, Sol-Attn2 compresses unselected blocks into a constant-size surrogate block to compensate for the attention mass discarded by sparsification. The surrogate and selected tokens share a common normalization, retaining global context with a fixed-size compensation representation. Combined with 8-bit quantization, Sol-Attn2 delivers **[xxx×]** end-to-end inference speedup on MiniMax H3 without training. After lightweight quantization-aware distillation (QAD), it enables higher sparsity and **[xxx×]** end-to-end speedup without quality degradation. *[E2E and quality results pending.]*

## Dense vs. Sol-Attn2

| Comparison | Methods | Status |
|---|---|---|
| Sample 01 | Dense / Sol-Attn2 | Coming soon |

## Method

Sol-Attn2 keeps token-level attention for selected blocks and represents the remaining tokens with compact, query-conditioned summaries. The summaries retain both probability mass and value contribution, so the two regions participate in one normalized output.

### Routing and complementary approximation

Dot products of 64-token Q/K block means select the exact set *Eᵢ* at the target sparsity. Each unselected key block is split into two content groups by its key projections along the mean-query direction. For a 32-token mean query *q̄ₐ*, let *𝒯ₐ* denote the tail groups, with sizes *n<sub>g</sub>* and mean keys and values *k̄<sub>g</sub>*, *v̄<sub>g</sub>*. With head dimension *d*,

$$
\begin{aligned}w_{ag}&=n_g e^{\bar q_a^\top\bar k_g/\sqrt d},\\ Z_a&=\sum_{g\in\mathcal T_a}w_{ag},\\ (\mu_a,\nu_a)&=\frac{1}{Z_a}\sum_{g\in\mathcal T_a}w_{ag}(\bar k_g,\bar v_g).\end{aligned}
$$

Group counts preserve mass, while *μₐ* and *νₐ* summarize the tail’s conditional key and value. Since the log-partition has query gradient *μₐ/√d*, a first-order expansion adjusts the mass for each query within the anchor:

$$
\begin{aligned}\delta_i&=\frac{(q_i-\bar q_a)^\top\mu_a}{\sqrt d},\\ \widehat Z_i&=Z_a e^{\delta_i},\qquad \widehat U_i=\widehat Z_i\nu_a.\end{aligned}
$$

Writing *zᵢⱼ = qᵢᵀkⱼ/√d*, exact and approximate contributions share the same denominator:

$$
o_i=\frac{\sum_{j\in E_i}e^{z_{ij}}v_j+\widehat U_i}{\sum_{j\in E_i}e^{z_{ij}}+\widehat Z_i}.
$$

Stage A1 optimizes only rank-64 QKVO LoRA, minimizing normalized squared error to the dense teacher’s layer outputs. Block selection and grouping are fixed during each backward pass; gradients flow through Q/K/V, the tail approximation, and joint normalization.

### How the outputs are combined

For query *i*, let *xᵢ ∈ ℝ<sup>D</sup>* be the input hidden state and *qᵢ, kⱼ, vⱼ ∈ ℝ<sup>d</sup>* the vectors for one attention head. Each method has its own selected key-token set *Eᵢ*. With *zᵢⱼ = qᵢᵀkⱼ/√d*, its sparse branch is

$$
\begin{aligned}Z_{E,i}&=\sum_{j\in E_i}e^{z_{ij}},\\ o_{s,i}&=\frac{\sum_{j\in E_i}e^{z_{ij}}v_j}{Z_{E,i}}.\end{aligned}
$$

The table gives the per-head output *oᵢ*. Heads are concatenated and passed through the usual O projection. Here *⊙* denotes elementwise multiplication.

| Method | Attention output | Auxiliary term / normalization |
|---|---|---|
| SLA | $o_i=o_{s,i}+W_L\ell_i+b_L$ | $\begin{aligned}a_{ij}&=\phi(q_i)^\top\phi(k_j),\\ \ell_i&=\frac{\sum_j a_{ij}v_j}{\sum_j a_{ij}+\epsilon}.\end{aligned}$ Linear attention over all tokens. *φ* is softmax over feature channels; *ε > 0* stabilizes the denominator. |
| VSA | $o_i=o_{s,i}+(W_{g,h}x_i)\odot c_i$ | $\begin{aligned}s_{ab}&=\frac{\bar q_a^\top\bar k_b}{\sqrt d},\\ \pi_{ab}&=\frac{e^{s_{ab}}}{\sum_c e^{s_{ac}}},\\ c_i&=\sum_b\pi_{b(i),b}\bar v_b.\end{aligned}$ Pooled attention over all blocks, broadcast from query tile *b(i)*. Bars denote block means. The same scores *s<sub>ab</sub>* also rank blocks for routing. |
| Sol-Attn2 | $o_i=\frac{Z_{E,i}o_{s,i}+\widehat Z_i\nu_{a(i)}}{Z_{E,i}+\widehat Z_i}$ | $\alpha_i=\frac{Z_{E,i}}{Z_{E,i}+\widehat Z_i}$ The tail contains only tokens outside *Eᵢ*. Its mass and value use query anchor *a(i)* from Eq. (2); *αᵢ* is the resulting sparse-branch weight. |

**Mixing parameters.** SLA learns an affine map, *W<sub>L</sub> ∈ ℝ<sup>d×d</sup>* and *b<sub>L</sub> ∈ ℝ<sup>d</sup>*, shared across heads. VSA uses *W<sub>g,h</sub> ∈ ℝ<sup>d×D</sup>*, the head-*h* slice of its full gate projection: *W<sub>g,h</sub>xᵢ* is an unconstrained coefficient for each output channel, with no bias or sigmoid. Sol-Attn2 obtains *αᵢ ∈ [0,1]* from attention masses, with no learned mixing gate.

SLA’s affine map and VSA’s gate start at zero, so each initially returns its own sparse branch. Sol-Attn2 includes complementary-tail compensation from initialization. All three A1 runs also train QKVO LoRA.

### FP8 quantization and ExpCast

The FP8 kernel quantizes Q and K with one E4M3 scale per 64-token block, and V with one scale per head. For a nonzero quantization block *X*,

$$
\begin{aligned}s_X&=\frac{\max_{ij}|X_{ij}|}{448},\\ X_8&=\operatorname{FP8}_{\mathrm{E4M3}}(X/s_X),\\ X&\approx s_X X_8.\end{aligned}
$$

QK and PV use FP8 tensor-core operands with FP32 accumulation. Scale factors restore magnitudes, and an additive bias on compressed keys preserves their mass at the query anchor.

Following VC-Attention, ExpCast maps base-2 scores directly to E4M3 probability bytes, replacing a per-score exponential followed by FP8 conversion. With a shared online pivot *m*, the unit-scale form used in this kernel is

$$
\begin{aligned}u&=(s-m)\log_2 e,\\ c(u)&=\operatorname{clip}_{[0,120]}\!\left(\operatorname{RNE}(8u+55.65)\right),\\ \widehat p(u)&=\operatorname{decode}_{\mathrm{E4M3}}(c(u))\approx 2^u.\end{aligned}
$$

RNE denotes rounding to nearest even. The same decoded weights enter the PV product and the denominator, keeping quantization consistent with the shared normalization.

## FP8 Kernel Performance

We compare Dense (cuDNN), Sol-Attn1, and the Sol-Attn2 FP8 kernel on GB200. At 90% sparsity, Sol-Attn2 is **1.54×, 2.06×, and 2.66× faster** than Sol-Attn1 at 32K, 64K, and 128K tokens, respectively.

![FP8 Kernel Performance](assets/kernel-10.svg)

GB200, 32 heads, head dimension 128; median full QKV-to-O latency, including input-dependent preparation. 90% sparsity. The FP8 kernel uses 64 compressed slots; A1 uses the Q32 whole-tail approximation described above.

| Tokens | Dense (cuDNN) / ms | Sol-Attn1 / ms | Sol-Attn2 (FP8) / ms | Speedup vs. Sol-Attn1 |
|---|---|---|---|---|
| 32K | 9.791 | 4.908 | 3.184 | 1.54× |
| 64K | 44.129 | 18.148 | 8.825 | 2.06× |
| 128K | 186.166 | 70.206 | 26.441 | 2.66× |

## Stage A1 Results

All three methods train for 200 steps across all 50 HyperFlow layers at fixed 90% token-pair sparsity, using the same training samples and noise sequence. Validation uses four held-out cases across eight noise intervals, evaluated at steps 0, 25, 50, 100, 150, and 200. Trainable parameter counts follow each method name.

![Stage A1 loss curves](assets/a1-curves.svg)

Left: stochastic training loss and trailing 15-step mean (log scale). Right: validation on fixed inputs (linear scale).

| Method (trainable parameters) | Step 0 | Step 200 |
|---|---:|---:|
| VSA (2,087.3M) | 0.05532 | 0.03907 |
| SLA (161.4M) | 0.06247 | 0.04686 |
| Sol-Attn2 (160.6M) | 0.05763 | 0.03821 |

Sol-Attn2 reduces validation loss from 0.05763 to **0.03821**, finishing 2.19% below VSA and 18.46% below SLA. It uses 1/13 of VSA’s trainable parameters and a similar count to SLA.

## More Samples

Further Dense and Sol-Attn2 comparisons with the same prompt and seed.

| Comparison | Methods | Status |
|---|---|---|
| Sample 02 | Dense / Sol-Attn2 | Coming soon |
| Sample 03 | Dense / Sol-Attn2 | Coming soon |
| Sample 04 | Dense / Sol-Attn2 | Coming soon |
