Flash Attention Residuals
4x faster inference/training vs. torch.compile naive attention residuals implementation
20% reduction in training memory (without activation checkpointing)*
*Benchmarked on H100. Dependent on problem size and setup.
Reference: https://arxiv.org/abs/2603.15031 (Kimi Team, MoonshotAI, 2026)
Credits:
Thanks to Mohamed Osman (https://github.com/spaghettiSystems) and Cartesia (https://github.com/cartesia-ai) for advising on and supporting the development of this project.
Install
pip install flash-attn-res
Usage
This package contains Triton kernels, triton_op wrappers compatible with torch.compile, and an experimental high-performance Block AttenRes autograd implementation.
See src and benchmarks folders.
Roadmap:
- Better autotuning set up
- Better benchmarks
- More robust autograd impl.
- Precision tuning
- Mixed FP16 and BF16 and store quantization scale
- Stochastic rounding
- CuTE, CUDA, and other DSLs implementation
Development Notes:
- Normalizing in phase 1 keeps outputs bounded (convex combination of values) so bf16 error doesn't scale with softmax flatness. Phase 2 computes in fp32, and the reduction algebra matches split-KV Flash Attention.
- Certain dimensions, especially NUM_QUERIES_PER_BLOCK, are small so semi-elementwise (B, T) kernel with static_range is better than doing tl.dot
- Kernel is memory bound and doing semi-elementwise allows for kernel fusion
- NUM_SOURCE_BLOCKS and NUM_QUERIES_PER_BLOCK should be autotuning keys, unlike with torch.compile, which allows for faster kernels
- Small NUM_QUERIES_PER_BLOCK so eviction_policy should be "evict_last"
Metadata
Release files for flash-attn-res 0.1.12
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| flash_attn_res-0.1.12.tar.gz | 1.3 MB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| flash_attn_res-0.1.12-py2.py3-none-any.whl | Python 2, Python 3 | none | any | Details |
Total release size: 1.3 MB
Release files / flash_attn_res-0.1.12.tar.gz
| Download URL | flash_attn_res-0.1.12.tar.gz |
|---|---|
| Size | 1.3 MB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
0f2b0bd6030a4d4af246229519dbaecdb1a2988c5113717b305baead904de9bf
|
|
BLAKE2b-256 checksum How to use checksums |
4c13f98fe7f640ff73a3cb642cf158fba6d6d918a5e1314d1fb0f6e90e6d817a
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/6.2.0 CPython/3.11.11
|
Release files / flash_attn_res-0.1.12-py2.py3-none-any.whl
| Download URL | flash_attn_res-0.1.12-py2.py3-none-any.whl |
|---|---|
| Size | 18.3 kB |
| Tags | Python 2 Python 3 |
|
SHA-256 checksum How to use checksums |
9e88c9dd69b74e4ac3d5958311b52b875ab6042a0e606eb9897090d7f988d714
|
|
BLAKE2b-256 checksum How to use checksums |
5f7c5547a0235b9e7c83ed46f94ed15a50acd276a05ab9938d494992ec0e2fbb
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/6.2.0 CPython/3.11.11
|