Skip to content

Make dropout-disabled FA2 libtorch ABI stable - #2688

Open
janeyx99 wants to merge 2 commits into
Dao-AILab:mainfrom
janeyx99:fa2-abi-stable
Open

janeyx99 wants to merge 2 commits into
Dao-AILab:mainfrom
janeyx99:fa2-abi-stable

Conversation

@janeyx99

@janeyx99 janeyx99 commented Jun 28, 2026 •

Copy link
Copy Markdown
Contributor

TL;DR: FA2 with dropout disabled is now ABI stable with torch 2.10+ and CPython 3.10+.

Here's how to think about each commit:

  • 9272e29 (this PR): Use TORCH_LIBRARY instead of pybind to enable building 1 wheel across multiple Python versions, thus enabling building FA2 in a CPython agnostic manner.
  • d5f02eb (this PR) Migrate to libtorch stable ABI. Since the APIs are only available 2.10+, I introduce branching based on build time torch version so that flash_api.cpp will be built if 2.10+, otherwise flash_api_unstable.cpp will be built. flash_api_unstable.cpp is identical to flash_api.cpp from the previous commit.

Test Plan:

Correctness

For the new stable dropout disabled build, I ran all tests that did not specify dropout of 0.17: FA2_TEST_NUM_GPUS=8 pytest tests/test_flash_attn.py -k "not 0.17" -q -p no:cacheprovider -n 8 --dist=load --tb=short -rfE
which yielded:

===================================== short test summary info =====================================
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[1-339-True-160-False-False-True-False-dtype0] - AssertionError: assert 0.00390625 <= ((8 * 0.0) + 0.0002)
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[1-339-True-160-False-False-True-True-dtype0] - AssertionError: assert 0.00390625 <= ((8 * 0.0) + 0.0002)
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[16-100000-False-96-True-False-True-False-dtype0] - AssertionError: assert 0.0020751953125 <= ((2 * 0.0009765625) + 1e-05)
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[16-100000-False-96-True-False-True-True-dtype0] - AssertionError: assert 0.0020751953125 <= ((2 * 0.0009765625) + 1e-05)
4 failed, 261717 passed, 152064 skipped in 353.06s (0:05:53)

Everything passed except these 4, which are pre-existing for me on main.

To parallelize across my GPUs, I used the following tests/conftest.py generated by Claude:

Details
import os

# Pin each pytest-xdist worker to a distinct GPU (round-robin) before any CUDA
# initialization, so `-n <num_gpus>` spreads the suite across all GPUs.
_worker = os.environ.get("PYTEST_XDIST_WORKER", "")
if _worker.startswith("gw"):
    _num_gpus = int(os.environ.get("FA2_TEST_NUM_GPUS", "8"))
    os.environ["CUDA_VISIBLE_DEVICES"] = str(int(_worker[2:]) % _num_gpus)

For the dropout enabled build of flash_api.cpp:, more tests, same results:
FA2_TEST_NUM_GPUS=8 python -m pytest tests/test_flash_attn.py -q -p no:cacheprovider -n 8 --dist=load --tb=short -rfE

============================================== short test summary info ==============================================
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[1-339-True-160-False-False-True-False-dtype0] - AssertionError: assert 0.00390625 <= ((8 * 0.0) + 0.0002)
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[1-339-True-160-False-False-True-True-dtype0] - AssertionError: assert 0.00390625 <= ((8 * 0.0) + 0.0002)
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[16-100000-False-96-True-False-True-False-dtype0] - AssertionError: assert 0.0020751953125 <= ((2 * 0.0009765625) + 1e-05)
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[16-100000-False-96-True-False-True-True-dtype0] - AssertionError: assert 0.0020751953125 <= ((2 * 0.0009765625) + 1e-05)
4 failed, 312357 passed, 196416 skipped in 705.62s (0:11:45)

And lastly, for a torch 2.9 build of flash_api_unstable.cpp, same pass:
FA2_TEST_NUM_GPUS=8 python -m pytest tests/test_flash_attn.py -q -p no:cacheprovider -n 8 --dist=load --tb=short -rfE

============================================== short test summary info ==============================================
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[1-339-True-160-False-False-True-False-dtype0] - AssertionError: assert 0.00390625 <= ((8 * 0.0) + 0.0002)
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[1-339-True-160-False-False-True-True-dtype0] - AssertionError: assert 0.00390625 <= ((8 * 0.0) + 0.0002)
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[16-100000-False-96-True-False-True-False-dtype0] - AssertionError: assert 0.0020751953125 <= ((2 * 0.0009765625) + 1e-05)
FAILED tests/test_flash_attn.py::test_flash_attn_splitkv[16-100000-False-96-True-False-True-True-dtype0] - AssertionError: assert 0.0020751953125 <= ((2 * 0.0009765625) + 1e-05)
4 failed, 312357 passed, 196416 skipped in 477.24s (0:07:57)

Perf

I ran the benchmark script for both before this commit (unstable .so) and after (stable .so) to verify that the difference is insignificant:

image

Stats and script are hidden below:

Details

On this branch (stable):

(fa2-pt210) ➜  flash-attention git:(fa2-abi-stable) ✗ CUDA_VISIBLE_DEVICES=1 taskset -c 0-32 python benchmarks/benchmark_flash_attention.py
CUDA_VISIBLE_DEVICES=1 taskset -c 0-32 python benchmarks/benchmark_flash_attention.py
### causal=False, headdim=64, batch_size=32, seqlen=512 ###
Flash2 fwd: 252.49 TFLOPs/s, bwd: 189.86 TFLOPs/s, fwd + bwd: 204.34 TFLOPs/s
Pytorch fwd: 43.18 TFLOPs/s, bwd: 51.42 TFLOPs/s, fwd + bwd: 48.76 TFLOPs/s
### causal=False, headdim=64, batch_size=16, seqlen=1024 ###
Flash2 fwd: 294.78 TFLOPs/s, bwd: 238.75 TFLOPs/s, fwd + bwd: 252.46 TFLOPs/s
Pytorch fwd: 50.61 TFLOPs/s, bwd: 57.11 TFLOPs/s, fwd + bwd: 55.09 TFLOPs/s
### causal=False, headdim=64, batch_size=8, seqlen=2048 ###
Flash2 fwd: 308.09 TFLOPs/s, bwd: 263.71 TFLOPs/s, fwd + bwd: 275.03 TFLOPs/s
Pytorch fwd: 50.10 TFLOPs/s, bwd: 65.18 TFLOPs/s, fwd + bwd: 60.02 TFLOPs/s
### causal=False, headdim=64, batch_size=4, seqlen=4096 ###
Flash2 fwd: 311.97 TFLOPs/s, bwd: 279.80 TFLOPs/s, fwd + bwd: 288.30 TFLOPs/s
Pytorch fwd: 49.06 TFLOPs/s, bwd: 68.12 TFLOPs/s, fwd + bwd: 61.32 TFLOPs/s
### causal=False, headdim=64, batch_size=2, seqlen=8192 ###
Flash2 fwd: 310.20 TFLOPs/s, bwd: 299.23 TFLOPs/s, fwd + bwd: 302.28 TFLOPs/s
Pytorch fwd: 44.90 TFLOPs/s, bwd: 69.67 TFLOPs/s, fwd + bwd: 60.19 TFLOPs/s
### causal=False, headdim=64, batch_size=1, seqlen=16384 ###
Flash2 fwd: 298.48 TFLOPs/s, bwd: 299.11 TFLOPs/s, fwd + bwd: 298.93 TFLOPs/s
Pytorch fwd: 64.72 TFLOPs/s, bwd: 69.52 TFLOPs/s, fwd + bwd: 68.08 TFLOPs/s
### causal=False, headdim=128, batch_size=32, seqlen=512 ###
Flash2 fwd: 299.40 TFLOPs/s, bwd: 190.57 TFLOPs/s, fwd + bwd: 212.65 TFLOPs/s
Pytorch fwd: 66.23 TFLOPs/s, bwd: 82.88 TFLOPs/s, fwd + bwd: 77.32 TFLOPs/s
### causal=False, headdim=128, batch_size=16, seqlen=1024 ###
Flash2 fwd: 340.76 TFLOPs/s, bwd: 238.05 TFLOPs/s, fwd + bwd: 260.48 TFLOPs/s
Pytorch fwd: 84.78 TFLOPs/s, bwd: 99.85 TFLOPs/s, fwd + bwd: 95.02 TFLOPs/s
### causal=False, headdim=128, batch_size=8, seqlen=2048 ###
Flash2 fwd: 370.21 TFLOPs/s, bwd: 273.29 TFLOPs/s, fwd + bwd: 295.38 TFLOPs/s
Pytorch fwd: 90.36 TFLOPs/s, bwd: 119.49 TFLOPs/s, fwd + bwd: 109.42 TFLOPs/s
### causal=False, headdim=128, batch_size=4, seqlen=4096 ###
Flash2 fwd: 377.48 TFLOPs/s, bwd: 289.66 TFLOPs/s, fwd + bwd: 310.29 TFLOPs/s
Pytorch fwd: 91.99 TFLOPs/s, bwd: 129.72 TFLOPs/s, fwd + bwd: 116.11 TFLOPs/s
### causal=False, headdim=128, batch_size=2, seqlen=8192 ###
Flash2 fwd: 357.99 TFLOPs/s, bwd: 308.66 TFLOPs/s, fwd + bwd: 321.31 TFLOPs/s
Pytorch fwd: 86.36 TFLOPs/s, bwd: 135.64 TFLOPs/s, fwd + bwd: 116.63 TFLOPs/s
### causal=False, headdim=128, batch_size=1, seqlen=16384 ###
Flash2 fwd: 366.66 TFLOPs/s, bwd: 305.34 TFLOPs/s, fwd + bwd: 320.66 TFLOPs/s
Pytorch fwd: 126.11 TFLOPs/s, bwd: 138.11 TFLOPs/s, fwd + bwd: 134.45 TFLOPs/s
### causal=True, headdim=64, batch_size=32, seqlen=512 ###
Flash2 fwd: 174.68 TFLOPs/s, bwd: 118.25 TFLOPs/s, fwd + bwd: 130.27 TFLOPs/s
Pytorch fwd: 14.71 TFLOPs/s, bwd: 25.80 TFLOPs/s, fwd + bwd: 21.23 TFLOPs/s
### causal=True, headdim=64, batch_size=16, seqlen=1024 ###
Flash2 fwd: 238.87 TFLOPs/s, bwd: 173.90 TFLOPs/s, fwd + bwd: 188.56 TFLOPs/s
Pytorch fwd: 16.40 TFLOPs/s, bwd: 28.56 TFLOPs/s, fwd + bwd: 23.57 TFLOPs/s
### causal=True, headdim=64, batch_size=8, seqlen=2048 ###
Flash2 fwd: 275.14 TFLOPs/s, bwd: 220.12 TFLOPs/s, fwd + bwd: 233.46 TFLOPs/s
Pytorch fwd: 15.94 TFLOPs/s, bwd: 32.57 TFLOPs/s, fwd + bwd: 25.09 TFLOPs/s
### causal=True, headdim=64, batch_size=4, seqlen=4096 ###
Flash2 fwd: 299.20 TFLOPs/s, bwd: 249.19 TFLOPs/s, fwd + bwd: 261.69 TFLOPs/s
Pytorch fwd: 13.98 TFLOPs/s, bwd: 34.04 TFLOPs/s, fwd + bwd: 24.14 TFLOPs/s
### causal=True, headdim=64, batch_size=2, seqlen=8192 ###
Flash2 fwd: 307.31 TFLOPs/s, bwd: 267.18 TFLOPs/s, fwd + bwd: 277.54 TFLOPs/s
Pytorch fwd: 13.05 TFLOPs/s, bwd: 34.84 TFLOPs/s, fwd + bwd: 23.59 TFLOPs/s
### causal=True, headdim=64, batch_size=1, seqlen=16384 ###
Flash2 fwd: 310.93 TFLOPs/s, bwd: 294.93 TFLOPs/s, fwd + bwd: 299.33 TFLOPs/s
Pytorch fwd: 16.53 TFLOPs/s, bwd: 34.75 TFLOPs/s, fwd + bwd: 26.43 TFLOPs/s
### causal=True, headdim=128, batch_size=32, seqlen=512 ###
Flash2 fwd: 191.00 TFLOPs/s, bwd: 129.37 TFLOPs/s, fwd + bwd: 142.51 TFLOPs/s
Pytorch fwd: 24.07 TFLOPs/s, bwd: 41.41 TFLOPs/s, fwd + bwd: 34.34 TFLOPs/s
### causal=True, headdim=128, batch_size=16, seqlen=1024 ###
Flash2 fwd: 262.47 TFLOPs/s, bwd: 190.22 TFLOPs/s, fwd + bwd: 206.46 TFLOPs/s
Pytorch fwd: 28.94 TFLOPs/s, bwd: 50.00 TFLOPs/s, fwd + bwd: 41.39 TFLOPs/s
### causal=True, headdim=128, batch_size=8, seqlen=2048 ###
Flash2 fwd: 303.89 TFLOPs/s, bwd: 236.78 TFLOPs/s, fwd + bwd: 252.73 TFLOPs/s
Pytorch fwd: 29.78 TFLOPs/s, bwd: 59.72 TFLOPs/s, fwd + bwd: 46.40 TFLOPs/s
### causal=True, headdim=128, batch_size=4, seqlen=4096 ###
Flash2 fwd: 325.09 TFLOPs/s, bwd: 265.49 TFLOPs/s, fwd + bwd: 280.17 TFLOPs/s
Pytorch fwd: 26.54 TFLOPs/s, bwd: 64.85 TFLOPs/s, fwd + bwd: 45.91 TFLOPs/s
### causal=True, headdim=128, batch_size=2, seqlen=8192 ###
Flash2 fwd: 341.13 TFLOPs/s, bwd: 289.05 TFLOPs/s, fwd + bwd: 302.24 TFLOPs/s
Pytorch fwd: 25.29 TFLOPs/s, bwd: 67.78 TFLOPs/s, fwd + bwd: 45.79 TFLOPs/s
### causal=True, headdim=128, batch_size=1, seqlen=16384 ###
Flash2 fwd: 335.62 TFLOPs/s, bwd: 319.17 TFLOPs/s, fwd + bwd: 323.71 TFLOPs/s
Pytorch fwd: 31.82 TFLOPs/s, bwd: 69.02 TFLOPs/s, fwd + bwd: 51.74 TFLOPs/s

On fully-yank-dropout (now main) (unstable, but with the same functionality):

(fa2-pt210) ➜  flash-attention git:(fully-yank-dropout) ✗ CUDA_VISIBLE_DEVICES=1 taskset -c 0-32 python benchmarks/benchmark_flash_attention.py
### causal=False, headdim=64, batch_size=32, seqlen=512 ###
Flash2 fwd: 252.53 TFLOPs/s, bwd: 189.25 TFLOPs/s, fwd + bwd: 203.84 TFLOPs/s
Pytorch fwd: 43.26 TFLOPs/s, bwd: 51.47 TFLOPs/s, fwd + bwd: 48.82 TFLOPs/s
### causal=False, headdim=64, batch_size=16, seqlen=1024 ###
Flash2 fwd: 295.27 TFLOPs/s, bwd: 239.06 TFLOPs/s, fwd + bwd: 252.81 TFLOPs/s
Pytorch fwd: 50.63 TFLOPs/s, bwd: 57.20 TFLOPs/s, fwd + bwd: 55.16 TFLOPs/s
### causal=False, headdim=64, batch_size=8, seqlen=2048 ###
Flash2 fwd: 310.09 TFLOPs/s, bwd: 257.43 TFLOPs/s, fwd + bwd: 270.56 TFLOPs/s
Pytorch fwd: 50.22 TFLOPs/s, bwd: 65.18 TFLOPs/s, fwd + bwd: 60.07 TFLOPs/s
### causal=False, headdim=64, batch_size=4, seqlen=4096 ###
Flash2 fwd: 310.32 TFLOPs/s, bwd: 279.80 TFLOPs/s, fwd + bwd: 287.89 TFLOPs/s
Pytorch fwd: 48.96 TFLOPs/s, bwd: 68.14 TFLOPs/s, fwd + bwd: 61.28 TFLOPs/s
### causal=False, headdim=64, batch_size=2, seqlen=8192 ###
Flash2 fwd: 306.62 TFLOPs/s, bwd: 299.81 TFLOPs/s, fwd + bwd: 301.72 TFLOPs/s
Pytorch fwd: 44.60 TFLOPs/s, bwd: 69.66 TFLOPs/s, fwd + bwd: 60.02 TFLOPs/s
### causal=False, headdim=64, batch_size=1, seqlen=16384 ###
Flash2 fwd: 302.94 TFLOPs/s, bwd: 297.38 TFLOPs/s, fwd + bwd: 298.95 TFLOPs/s
Pytorch fwd: 64.70 TFLOPs/s, bwd: 69.52 TFLOPs/s, fwd + bwd: 68.07 TFLOPs/s
### causal=False, headdim=128, batch_size=32, seqlen=512 ###
Flash2 fwd: 299.58 TFLOPs/s, bwd: 190.48 TFLOPs/s, fwd + bwd: 212.60 TFLOPs/s
Pytorch fwd: 65.88 TFLOPs/s, bwd: 82.19 TFLOPs/s, fwd + bwd: 76.76 TFLOPs/s
### causal=False, headdim=128, batch_size=16, seqlen=1024 ###
Flash2 fwd: 341.31 TFLOPs/s, bwd: 238.52 TFLOPs/s, fwd + bwd: 260.98 TFLOPs/s
Pytorch fwd: 84.93 TFLOPs/s, bwd: 100.11 TFLOPs/s, fwd + bwd: 95.24 TFLOPs/s
### causal=False, headdim=128, batch_size=8, seqlen=2048 ###
Flash2 fwd: 365.52 TFLOPs/s, bwd: 271.10 TFLOPs/s, fwd + bwd: 292.70 TFLOPs/s
Pytorch fwd: 90.25 TFLOPs/s, bwd: 119.37 TFLOPs/s, fwd + bwd: 109.29 TFLOPs/s
### causal=False, headdim=128, batch_size=4, seqlen=4096 ###
Flash2 fwd: 374.05 TFLOPs/s, bwd: 292.16 TFLOPs/s, fwd + bwd: 311.65 TFLOPs/s
Pytorch fwd: 92.42 TFLOPs/s, bwd: 129.73 TFLOPs/s, fwd + bwd: 116.31 TFLOPs/s
### causal=False, headdim=128, batch_size=2, seqlen=8192 ###
Flash2 fwd: 352.59 TFLOPs/s, bwd: 307.71 TFLOPs/s, fwd + bwd: 319.33 TFLOPs/s
Pytorch fwd: 86.35 TFLOPs/s, bwd: 135.67 TFLOPs/s, fwd + bwd: 116.64 TFLOPs/s
### causal=False, headdim=128, batch_size=1, seqlen=16384 ###
Flash2 fwd: 367.87 TFLOPs/s, bwd: 303.27 TFLOPs/s, fwd + bwd: 319.29 TFLOPs/s
Pytorch fwd: 126.01 TFLOPs/s, bwd: 138.15 TFLOPs/s, fwd + bwd: 134.45 TFLOPs/s
### causal=True, headdim=64, batch_size=32, seqlen=512 ###
Flash2 fwd: 175.15 TFLOPs/s, bwd: 118.33 TFLOPs/s, fwd + bwd: 130.42 TFLOPs/s
Pytorch fwd: 14.71 TFLOPs/s, bwd: 25.81 TFLOPs/s, fwd + bwd: 21.23 TFLOPs/s
### causal=True, headdim=64, batch_size=16, seqlen=1024 ###
Flash2 fwd: 238.76 TFLOPs/s, bwd: 173.09 TFLOPs/s, fwd + bwd: 187.85 TFLOPs/s
Pytorch fwd: 16.36 TFLOPs/s, bwd: 28.60 TFLOPs/s, fwd + bwd: 23.57 TFLOPs/s
### causal=True, headdim=64, batch_size=8, seqlen=2048 ###
Flash2 fwd: 275.99 TFLOPs/s, bwd: 218.35 TFLOPs/s, fwd + bwd: 232.21 TFLOPs/s
Pytorch fwd: 15.88 TFLOPs/s, bwd: 32.57 TFLOPs/s, fwd + bwd: 25.05 TFLOPs/s
### causal=True, headdim=64, batch_size=4, seqlen=4096 ###
Flash2 fwd: 299.40 TFLOPs/s, bwd: 243.87 TFLOPs/s, fwd + bwd: 257.52 TFLOPs/s
Pytorch fwd: 13.98 TFLOPs/s, bwd: 34.05 TFLOPs/s, fwd + bwd: 24.14 TFLOPs/s
### causal=True, headdim=64, batch_size=2, seqlen=8192 ###
Flash2 fwd: 300.37 TFLOPs/s, bwd: 273.32 TFLOPs/s, fwd + bwd: 280.54 TFLOPs/s
Pytorch fwd: 13.06 TFLOPs/s, bwd: 34.83 TFLOPs/s, fwd + bwd: 23.59 TFLOPs/s
### causal=True, headdim=64, batch_size=1, seqlen=16384 ###
Flash2 fwd: 301.54 TFLOPs/s, bwd: 294.52 TFLOPs/s, fwd + bwd: 296.49 TFLOPs/s
Pytorch fwd: 16.54 TFLOPs/s, bwd: 34.74 TFLOPs/s, fwd + bwd: 26.43 TFLOPs/s
### causal=True, headdim=128, batch_size=32, seqlen=512 ###
Flash2 fwd: 192.48 TFLOPs/s, bwd: 129.69 TFLOPs/s, fwd + bwd: 143.02 TFLOPs/s
Pytorch fwd: 24.16 TFLOPs/s, bwd: 41.51 TFLOPs/s, fwd + bwd: 34.44 TFLOPs/s
### causal=True, headdim=128, batch_size=16, seqlen=1024 ###
Flash2 fwd: 264.67 TFLOPs/s, bwd: 190.91 TFLOPs/s, fwd + bwd: 207.43 TFLOPs/s
Pytorch fwd: 28.96 TFLOPs/s, bwd: 49.92 TFLOPs/s, fwd + bwd: 41.37 TFLOPs/s
### causal=True, headdim=128, batch_size=8, seqlen=2048 ###
Flash2 fwd: 302.26 TFLOPs/s, bwd: 235.89 TFLOPs/s, fwd + bwd: 251.68 TFLOPs/s
Pytorch fwd: 29.76 TFLOPs/s, bwd: 59.73 TFLOPs/s, fwd + bwd: 46.38 TFLOPs/s
### causal=True, headdim=128, batch_size=4, seqlen=4096 ###
Flash2 fwd: 330.48 TFLOPs/s, bwd: 260.23 TFLOPs/s, fwd + bwd: 277.05 TFLOPs/s
Pytorch fwd: 26.77 TFLOPs/s, bwd: 64.89 TFLOPs/s, fwd + bwd: 46.13 TFLOPs/s
### causal=True, headdim=128, batch_size=2, seqlen=8192 ###
Flash2 fwd: 345.91 TFLOPs/s, bwd: 287.04 TFLOPs/s, fwd + bwd: 301.71 TFLOPs/s
Pytorch fwd: 25.33 TFLOPs/s, bwd: 67.78 TFLOPs/s, fwd + bwd: 45.83 TFLOPs/s
### causal=True, headdim=128, batch_size=1, seqlen=16384 ###
Flash2 fwd: 331.04 TFLOPs/s, bwd: 317.95 TFLOPs/s, fwd + bwd: 321.59 TFLOPs/s
Pytorch fwd: 31.79 TFLOPs/s, bwd: 69.02 TFLOPs/s, fwd + bwd: 51.72 TFLOPs/s

Script to compare them:

#!/usr/bin/env python3
"""Compare TFLOP/s between two benchmark_flash_attention.py runs.

Parses the stdout of two runs (e.g. before/after a change) and prints a
per-config fwd / bwd / fwd+bwd delta table for one method (default Flash2).

Usage:
    # baseline build
    CUDA_VISIBLE_DEVICES=1 python benchmarks/benchmark_flash_attention.py | tee before.txt
    # ... rebuild / switch branch ...
    CUDA_VISIBLE_DEVICES=1 python benchmarks/benchmark_flash_attention.py | tee after.txt
    python agent_space/compare_bench.py before.txt after.txt
    python agent_space/compare_bench.py before.txt after.txt --method Pytorch
"""
import argparse
import math
import re
import sys

GREEN, RED, RESET = "\033[32m", "\033[31m", "\033[0m"

CONFIG_RE = re.compile(r"^###\s*(.+?)\s*###\s*$")
ROW_RE = re.compile(
    r"^(\S+)\s+fwd:\s*([^\s,]+)\s*TFLOPs/s,\s*"
    r"bwd:\s*([^\s,]+)\s*TFLOPs/s,\s*"
    r"fwd \+ bwd:\s*([^\s,]+)\s*TFLOPs/s"
)


def parse(path):
    """path -> {config_str: {method: (fwd, bwd, fwd_bwd)}}, preserving order."""
    out = {}
    cfg = None
    with open(path) as f:
        for line in f:
            m = CONFIG_RE.match(line)
            if m:
                cfg = m.group(1)
                out.setdefault(cfg, {})
                continue
            m = ROW_RE.match(line)
            if m and cfg is not None:
                method = m.group(1)
                fwd, bwd, fb = (float(x) for x in m.groups()[1:])
                out[cfg][method] = (fwd, bwd, fb)
    return out


def short_cfg(s):
    """'causal=False, headdim=64, batch_size=32, seqlen=512' -> compact label."""
    d = dict(kv.split("=") for kv in s.replace(", ", ",").split(",") if "=" in kv)
    return (f"c={d.get('causal','?'):<5} hd={d.get('headdim','?'):<3} "
            f"bs={d.get('batch_size','?'):<3} sl={d.get('seqlen','?'):<5}")


def pct(before, after):
    if before == 0 or math.isnan(before) or math.isnan(after):
        return float("nan")
    return (after / before - 1.0) * 100.0


def fmt_pct(p):
    s = "n/a" if math.isnan(p) else f"{p:+.1f}%"
    return f"{s:>8}"


def colorize(s, p, enable):
    """Wrap a (pre-padded) string green if p>=0, red if p<0; pass-through if disabled."""
    if not enable or math.isnan(p):
        return s
    return f"{GREEN if p >= 0 else RED}{s}{RESET}"


def main():
    ap = argparse.ArgumentParser(description=__doc__,
                                 formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("before")
    ap.add_argument("after")
    ap.add_argument("--method", default="Flash2", help="method to compare (default: Flash2)")
    ap.add_argument("--threshold", type=float, default=2.0,
                    help="flag regressions worse than this %% (default: 2.0)")
    ap.add_argument("--color", choices=["auto", "always", "never"], default="auto",
                    help="colorize deltas (default: auto = only on a TTY)")
    args = ap.parse_args()

    use_color = args.color == "always" or (args.color == "auto" and sys.stdout.isatty())
    before, after = parse(args.before), parse(args.after)
    method = args.method

    # Keep the order of the baseline file, then append any after-only configs.
    configs = list(before) + [c for c in after if c not in before]

    g = "fwd TF/s".center(24) + "  " + "bwd TF/s".center(24) + "  " + "fwd+bwd TF/s".center(24)
    print(f"\nmethod = {method}   ({args.before}  ->  {args.after})\n")
    print(f"{'':<28}{g}")
    sub = f"{'before':>7}{'after':>8}{'Δ':>8} "
    print(f"{'config':<28}{sub} {sub} {sub}")
    print("-" * (28 + 3 * 24 + 4))

    # ratio products for geomean per metric; worst regression tracker
    log_sum = [0.0, 0.0, 0.0]
    log_n = [0, 0, 0]
    worst = (float("inf"), None, None)  # (delta%, config, metric)

    for cfg in configs:
        b = before.get(cfg, {}).get(method)
        a = after.get(cfg, {}).get(method)
        if b is None or a is None:
            miss = args.before if b is None else args.after
            print(f"{short_cfg(cfg):<28}  (missing in {miss})")
            continue
        cells = ""
        for i, name in enumerate(("fwd", "bwd", "fwd+bwd")):
            p = pct(b[i], a[i])
            cells += f"{b[i]:>7.1f}{a[i]:>8.1f}{colorize(fmt_pct(p), p, use_color)} "
            if not math.isnan(p):
                log_sum[i] += math.log(a[i] / b[i])
                log_n[i] += 1
                if p < worst[0]:
                    worst = (p, cfg, name)
        print(f"{short_cfg(cfg):<28}{cells}")

    print("-" * (28 + 3 * 24 + 4))
    geo = ""
    for i in range(3):
        gp = (math.exp(log_sum[i] / log_n[i]) - 1) * 100 if log_n[i] else float("nan")
        geo += f"{'':>15}{colorize(fmt_pct(gp), gp, use_color)} "
    print(f"{'geomean Δ':<28}{geo}")

    if worst[1] is not None and worst[0] < -args.threshold:
        msg = (f"⚠  worst regression: {worst[0]:+.1f}% on {worst[2]} "
               f"@ {short_cfg(worst[1]).strip()}")
        print("\n" + (f"{RED}{msg}{RESET}" if use_color else msg))
    elif worst[1] is not None:
        msg = (f"✓  no regression worse than {args.threshold:.1f}% "
               f"(worst: {worst[0]:+.1f}% on {worst[2]} @ {short_cfg(worst[1]).strip()})")
        print("\n" + (f"{GREEN}{msg}{RESET}" if use_color else msg))


if __name__ == "__main__":
    main()

No unstable APIs

68 stable! no unstable in the dropout disabled .so.

(fa2-pt210) ➜  flash-attention git:(fa2-abi-stable) ✗ torch-abi-audit ./flash_attn_2_cuda.abi3.so
Package: flash_attn_2_cuda.abi3.so
  Root: /home/janeyx/repos/flash-attention
  Torch ABI:   STABLE
  CPython ABI: n/a
  Extensions:  0
  Bundled libs: 1
  -- bundled libs --
    [STABLE  ] [abi3-tagged-no-capi   ] flash_attn_2_cuda.abi3.so  (stable_shim=68, unstable=0)

@janeyx99
janeyx99 force-pushed the fa2-abi-stable branch 3 times, most recently from 560efae to a61b003 Compare June 28, 2026 22:40
Comment thread csrc/flash_attn/flash_api.cpp Outdated

std::vector<at::Tensor>
mha_fwd(at::Tensor &q, // batch_size x seqlen_q x num_heads x round_multiple(head_size, 8)
mha_fwd(at::Tensor q, // batch_size x seqlen_q x num_heads x round_multiple(head_size, 8)

@janeyx99 janeyx99 Jun 29, 2026 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

For the pybind -> TORCH_LIBRARY commit: q is made pass-by-value as it is readonly (underlying data does not change) even while q gets reassigned in the method. Some references are dropped because registering through torch does not support non-const optional references. Dropping references here does not really matter because at::Tensor is just a pointer already.

std::optional<at::Tensor> &out_, // batch_size x seqlen_q x num_heads x round_multiple(head_size, 8)
std::optional<at::Tensor> &alibi_slopes_, // num_heads or batch_size x num_heads
const float p_dropout,
const float softmax_scale,

@janeyx99 janeyx99 Jun 29, 2026 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

for the pybind -> TORCH_LIBRARY commit: torch expects specific argument types, e.g., https://github.com/pytorch/pytorch/blob/12a9ea264bf805a66cd87e19e767ab23c2f59fef/aten/src/ATen/core/boxing/impl/make_boxed_from_unboxed_functor.h#L213. Namely, in this file:

  • int -> int64_t
  • float -> double
  • std::optional & -> const std::optionalat::Tensor &

Comment thread csrc/flash_attn/flash_api.cpp Outdated
m.def("bwd", &FLASH_NAMESPACE::mha_bwd, "Backward pass");
m.def("varlen_bwd", &FLASH_NAMESPACE::mha_varlen_bwd, "Backward pass (variable length)");
m.def("fwd_kvcache", &FLASH_NAMESPACE::mha_fwd_kvcache, "Forward pass, with KV-cache");
TORCH_LIBRARY(flash_attn_2, m) {

@janeyx99 janeyx99 Jun 29, 2026 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

for pybind -> TORCH_LIBRARY commit: Schemas were determined by actual data flow (is the actual tensor mutated or not). For the ops that the vllm fork also use, we share almost everything. The two differences are:

  • we mark q as not mutable since it's not actually written into
  • we mark k/vcache as mutated because they are appended to

Comment thread setup.py
# Tag with py_limited_api only for the CUDA build which is CPython agnostic by avoiding pybind.
options={"bdist_wheel": {"py_limited_api": "cp39"}} if ext_modules and not IS_ROCM else {},
python_requires=">=3.9",
options={"bdist_wheel": {"py_limited_api": "cp310"}} if ext_modules and not IS_ROCM else {},

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

torch 2.10 has min python 3.10

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant