Everything below was run on one machine; the exact software pins live in
VERIFICATION.md and inside each run's meta_start.json. Short version:
RTX PRO 6000 Blackwell WS (sm_120, 96 GB GDDR7), Ubuntu 24.04, Python 3.11/3.12,
PyTorch 2.11.0+cu128, CUDA 12.8, FlashInfer 0.6.13, nvidia-cutlass-dsl==4.2.0.
python3 -m venv venv && . venv/bin/activate
pip install "torch==2.11.*" --index-url https://download.pytorch.org/whl/cu128
pip install nvidia-cutlass-dsl==4.2.0 flashinfer-python==0.6.13 \
transformers numpyThe harnesses assume model weights are available through the usual Hugging Face
cache (HF_HOME). Absolute /home/USER/... paths inside the scripts are
placeholders — point them at your checkout.
The public FA2–FA4 wheels do not cover sm_120, so FlashAttention-4 is built from source. Use the companion repository, which is the upstream tree at the pinned base with the window already applied:
git clone https://github.com/ATFlash/ATFlash-Kernel-FA4
cd ATFlash-Kernel-FA4
pip install -e . # CuTeDSL path: python-only, no C++ compilation at installThe raw diffs live in that repository's patches/ directory; its
patches/FA4_PATCH_RECONCILIATION.md maps each patch to its measured number
(the paper's "+296/−13", the measured binary's "+308/−13", and the full port).
Its git history applies them one commit per step, so git diff between
commits shows exactly what the window adds.
The window is opt-in: without the atflash_window_zhi argument the code path
is stock.
The released 0.6.13 wheel is patched as a shadow copy — upstream is not rebuilt:
pip download flashinfer-python==0.6.13 --no-deps -d /tmp/fi
mkdir -p shadow && cd shadow && unzip -q /tmp/fi/*.whl
patch -p1 -d flashinfer < /path/to/kernels/flashinfer_port/decode.cuh.diff
patch -p1 -d flashinfer < /path/to/kernels/flashinfer_port/utils.cuh.diff
export PYTHONPATH=$PWD:$PYTHONPATH
export FLASHINFER_WORKSPACE_BASE=$PWD/cachekernels/flashinfer_port/fi_window_helpers.py shows the JIT-module builder and
the vec8 window-table snap this port consumes.
python window_tables/c1_window_precheck.py # CPU only, no weightsprints the argsort separation gate and the zhi tables for several k, with
hashes. window_tables/w_meta_from_hf.py derives tables for any HF model from
its config (inv_freq); the w_meta_*.json files are the exact tables the
published runs used.
python gates/b0_dualkernel_smoke.py # window OFF == stock, bit-identical, both ports
python gates/c0_chunk_neutrality.py # chunked prefill == single-shot, OFF and ON
python gates/g1_bitexact.py # FlashInfer decode OFF == stock (needs the shadow copy)
python gates/g2_correctness.py # fp32 correctness of the windowed pathg2 compares against atflash/reference.py; its captured-activation inputs
are large and not bundled — it regenerates or expects them locally.
# prefill microbenchmark (Qwen2.5-7B-1M shapes, 8K..512K, 4 systems)
python microbench/phasec_7b_prefill_mb.py
# end-to-end whole-request, window ON/OFF in the same binary
python e2e_pipeline/atflash_e2e_fafi.py --model qwen7b1m --bench longbench \
--lb-long-only --ctx-cap 262144 --prefill-chunk 65536 --no-cudnn --timing-onlye2e_pipeline/README.md explains the pool builders and the flags; every
recorded run's exact configuration is in its meta_start.json, and the raw
summaries behind the tables are in results_raw/summaries/.