Conversation
560efae to
a61b003
Compare
|
|
||
| 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) |
There was a problem hiding this comment.
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, |
There was a problem hiding this comment.
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 &
| 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) { |
There was a problem hiding this comment.
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
| # 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 {}, |
There was a problem hiding this comment.
torch 2.10 has min python 3.10
6a78eae to
05886b5
Compare
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 thatflash_api.cppwill be built if 2.10+, otherwiseflash_api_unstable.cppwill be built.flash_api_unstable.cppis identical toflash_api.cppfrom 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 -rfEwhich yielded:
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
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 -rfEAnd 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 -rfEPerf
I ran the benchmark script for both before this commit (unstable .so) and after (stable .so) to verify that the difference is insignificant:
Stats and script are hidden below:
Details
On this branch (stable):
On fully-yank-dropout (now main) (unstable, but with the same functionality):
Script to compare them:
No unstable APIs
68 stable! no unstable in the dropout disabled .so.