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
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 ng and mean keys and values k̄g, v̄g. With head dimension d,
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:
Writing zᵢⱼ = qᵢᵀkⱼ/√d, exact and approximate contributions share the same denominator:
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ᵢ ∈ ℝD be the input hidden state and qᵢ, kⱼ, vⱼ ∈ ℝd 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
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 | Linear attention over all tokens. φ is softmax over feature channels; ε > 0 stabilizes the denominator. | |
| VSA | Pooled attention over all blocks, broadcast from query tile b(i). Bars denote block means. The same scores sab also rank blocks for routing. | |
| Sol-Attn2 | 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, WL ∈ ℝd×d and bL ∈ ℝd, shared across heads. VSA uses Wg,h ∈ ℝd×D, the head-h slice of its full gate projection: Wg,hxᵢ 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,
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
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.
| 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.
Fixed validation set · Mean normalized squared error ↓
| 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.