◐ Off-By-One · answer catalog

python-flash-attention-online-softmax-block-recompute

1 answer(s)pythonpython3

python-flash-attention-online-softmax-block-recompute

📦 Source in repository (JSON)

Answer

All checks pass. Here's the solution.

Summary

I implemented and verified a memory-efficient exact attention kernel in pure NumPy.

Root cause: The naive softmax(QKᵀ)V materializes an N×N score/probability matrix, giving O(N²) peak memory, and the standard backward needs P again, forcing it to be stored.

The fix (3 coupled parts): 1. Online softmax forward — iterate key blocks keeping running max m, denominator l, and output accumulator; rescale l and acc by e^{m_old−m_new} when the max grows. Save only the per-row logsumexp L = m + log l. 2. Recompute backward — given (Q,K,V,O,dO,L), recompute P = exp(S−L) block-by-block instead of storing it, so no N×N tensor ever exists. Gradients use D_i = dO_i·O_i, dS = P·(dP − D_i), dQ = dS·K/√D, dK = dSᵀ·Q/√D, dV = Pᵀ·dO. 3. Block causal masking — skip key blocks right of the query block, mask only the diagonal block.

Files: - ~/flash_attn.py — implementation (flash_attn_forward, flash_attn_backward, reference_attention) - ~/test_flash_attn.py — verification suite - ~/SOLUTION.md — full write-up with code, math, and results

Verification results: - Forward vs reference: max rel err ≈2e-13 across causal/non-causal, non-power-of-two N, and D not divisible by block size. - Backward vs independent full-matrix backward: ≈9e-12. - Backward vs central finite differences: ≈4e-8 — all within the required 1e-5. - Memory: with block=32, max block temporary stays at exactly 1024 elements as N grows 64→1024 (temp/N² → 0). tracemalloc shows peak scaling linearly: flash 4.2 MB at N=2048 vs naive 105 MB; the naive path allocates N².

Evidence & signatures

# Evidence
- Problem class: python-flash-attention-online-softmax-block-recompute
- Model: openrouter/deepseek/deepseek-v4.1-flash
- Solved: 2026-10-02T16:18:50.795Z
- Verification: solution produced by pi in sandbox; see signatures.json
{"description": "Implement a memory-efficient exact attention kernel in pure NumPy/PyTorch-style tensors: forward pass with online (streaming) softmax over key blocks so the full S = QK^T matrix is never materialized, plus a backward pass that recomputes block statistics instead of storing the attention matrix. Support causal masking, head dimension not divisible by the block size, and sequence lengths that are not powers of two. Validate gradients against torch.autograd (or finite differences) to within 1e-5 relative error and prove peak intermediate memory is O(block_size^2) rather than O(seq_len^2).", "environment": "python3", "language": "python", "model": "openrouter/deepseek/deepseek-v4.1-flash", "problem_class": "python-flash-attention-online-softmax-block-recompute", "provider": "openrouter", "solved_at": "2026-10-02T16:18:50.795Z", "version": "3.11"}
Generated from the verified corpus · MIT licensedBack to the catalog