Skip to main content

FlexHeadFA: Flash Attention with flexible head dimensions

Project description

FlexHeadFA

This repository is a fork of the FlashAttention main repo. It extends the official implementation to support FlashAttention with flexible head dimensions.

All configurations in FlashAttention-2 are supported. Besides, we have supported:

  • FlashAttention-2 with QKHeadDim=32, VHeadDim=64
  • FlashAttention-2 with QKHeadDim=64, VHeadDim=128
  • FlashAttention-2 with QKHeadDim=96, VHeadDim=192
  • FlashAttention-2 with QKHeadDim=128, VHeadDim=256
  • FlashAttention-2 with QKHeadDim=192, VHeadDim=128
  • FLashAttention-2 with not equal num_heads_k and num_heads_v, such as (num_heads_q, num_heads_k, num_heads_v) = (32, 4, 16)

For headdim not supported, you can use the autotuner to generate the implementation. Details are in autotuner.md.

Feel free to tell us what else you need. We might support it soon. :)

Installation

The requirements is the same as FlashAttention-2

To install:

pip install flex-head-fa --no-build-isolation

Alternatively you can compile from source:

python setup.py install

The usage remains the same as FlashAttention-2. You only need to replace flash_attn with flex_head_fa, as shown below:

from flex_head_fa import flash_attn_func, flash_attn_with_kvcache

We are also developing FlexHeadFA based on the lastest FLashAttention-3. Currently, besides all configurations in FlashAttention-3, we also support

  • FlashAttention-3 with QKHeadDim=32, VHeadDim=64

  • FlashAttention-3 forward + FlashAttention-2 backward with QKHeadDim=128, VHeadDim=256 (FlashAttention-3 backward is under development)

Try it with:

cd hopper
python setup.py install

Usage:

from flash_attn_interface import flash_attn_func # FlashAttention-3 forward+backward
from flash_attn_interface import flash_attn_f3b2_func as flash_attn_func # FlashAttention-3 forward + FlashAttention-2 backward 

Performance of FlexHeadFA

We test the performance speedup compare to padding qk&v hidden_dim on A100.

We display FlexHeadFA speedup using these parameters:

  • (qk dim, v_dim): (32,64), (64,128), (128,256); qk hidden dimension 2048 (i.e. 64, 32 or 16 heads).
  • Sequence length 512, 1k, 2k, 4k, 8k, 16k.
  • Batch size set to 16k / seqlen.

Speedup

Custom-flash-attn

When you encounter issues

This new release of FlexHeadFA has been tested on several GPT-style models, mostly on A100 GPUs.

If you encounter bugs, please open a GitHub Issue!

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

flex_head_fa-0.1.2.tar.gz (2.8 MB view details)

Uploaded Source

File details

Details for the file flex_head_fa-0.1.2.tar.gz.

File metadata

  • Download URL: flex_head_fa-0.1.2.tar.gz
  • Upload date:
  • Size: 2.8 MB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.0.1 CPython/3.12.0

File hashes

Hashes for flex_head_fa-0.1.2.tar.gz
Algorithm Hash digest
SHA256 a52717576882a5765f5f52ac86b8bd0a251d8b0f104438ae429c2144463662b7
MD5 8c48351d472ee71579533103dc5beac6
BLAKE2b-256 980e63d9e6318c4b09ba52653b9c51645b1922908028bb5784808ab8400b90e7

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page