Skip to content

[None][feat] Add VisualGen FA4 autotuning and patched FA4 wheel build - #19668

Closed
chang-l wants to merge 5 commits into
NVIDIA:mainfrom
chang-l:demo/visualgen-fa4-autotune-20260928
Closed

chang-l wants to merge 5 commits into
NVIDIA:mainfrom
chang-l:demo/visualgen-fa4-autotune-20260928

Conversation

@chang-l

@chang-l chang-l commented Sep 28, 2026 •

Copy link
Copy Markdown
Collaborator

FA4's CTA count and exp2 emulation policy can depend on the GPU and attention shape. This opt-in draft tunes those choices through VisualGen's existing warmup AutoTuner, caches the result for inference, and retains FA4's default heuristic for untuned inputs.

The tuner covers dense noncausal FP16/BF16 MHA with head dimension 128 on SM100/SM103. It compares 1CTA and eligible 2CTA with a bounded set of exp2 policies, preserves output and float32 LSE, and includes automatic split-KV as a baseline candidate. Exact shapes, dtype, strides, device, dependency/build identity and dispatch policy separate cache entries. Enable with TLLM_VISUAL_GEN_FA4_AUTOTUNE=1. FA_DISABLE_2CTA=1 and FA4's CUDA 12 guard remain authoritative, complementing #19292.

FA4 b19 lacks the required per-call controls. 3rdparty/patches/flash_attn_4_b19.patch adds CTA/exp2 arguments, validates eligible inputs and includes exp2 policy in the compilation key without changing shared tuning tables. Calls omitting the new arguments retain their original behavior.

The build produces a separate patched FA4 dependency wheel:

  • CMake FetchContent fetches upstream revision 940cd9680f3315f2f06b43ab5bea2c2cf2d96806 and applies the patch through 3rdparty/prepare_fa4.py. The revision corresponds to b19; all 50 upstream package Python files match the published wheel before patching.
  • Revision, patch and package-input digest checks reject mismatched or stale sources. Wheel staging repeats validation even when CMake reuses populated sources. Version checks require the bootstrap and runtime pins to match the validated upstream and patched versions.
  • scripts/build_wheel.py uses FA4's upstream setuptools backend to build flash_attn_4-4.0.0b19+trtllm.1-py3-none-any.whl, installs it before generating TRT-LLM Python stubs, and emits it beside the TRT-LLM wheel. The package retains flash_attn.cute, includes upstream license/author files and adds build provenance for autotuner cache separation. Kernels still JIT-compile on the target GPU.
  • The TRT-LLM wheel requires exactly flash-attn-4==4.0.0b19+trtllm.1 and contains no FA4 package files. Stock b19 cannot satisfy that dependency. Source bootstrap requirements retain upstream b19; reinstalling them accepts the installed patched local version.
  • Release containers resolve the companion wheel with --find-links; Jenkins artifact archives and uploads carry both wheels. Test paths using --no-deps install both explicitly, and precompiled installation checks receive the companion artifact location. Python examples and trtllm-serve use the installed patched package without a runtime patch step.
  • Existing FA4 import paths are preserved, including Sol-Attn and DFlash. There is no change to sana-sol-attn.patch, its vendor lock, or vendored Sol-Attn sources in the final PR diff.

See build and upgrade notes and the standalone tuning demo. No upstream FA4 version upgrade is included. An upstream register-allocation fix alone does not replace this per-call API.

Validation:

  • 66 unique pytest cases passed: 11 CPU packaging/staging checks, 13 precompiled-install checks, and 42 FA4 checks in an SM103 GPU environment (including 12 tuner/cache policy checks).
  • The numerical matrix covers FP16/BF16, self- and cross-attention, ragged lengths, contiguous/head-sliced/packed-QKV layouts, output plus float32 LSE against an FP32 reference, split-KV behavior, and CUDA graph replay. Every eligible SM103 candidate passed.
  • Real warmup profiling, cache persistence and reload in a fresh process, unknown-shape fallback, and graph capture inside the autotune context passed. Reload profiles no tactics. Testing found and fixed graph capture discarding the cached winner, with CPU and CUDA regression coverage.
  • The standalone benchmark ran with the actual patched dependency at sequence lengths 4096, 14040 and 65520. Its FA4 import order now respects the backend's CUTLASS compatibility shims.
  • Wan2.2 A14B BF16, 480x832x33, 4-step CUDA-graph smoke passed through the full pipeline, public Python API and trtllm-serve HTTP endpoint. Default and tuned outputs, and the outputs across these entry points, are bitwise identical (LPIPS 0). Actual pipeline warmup populates FA4 cache entries; steady inference does not profile new tactics.
  • Packaging tests build a TRT-LLM fixture through the real setup.py and resolve the separate patched FA4 dependency with --find-links; stock b19 cannot satisfy the exact runtime pin. Source, patch and version mismatch rejection and incremental staging are covered. The actual CMake FetchContent/patch and upstream FA4 wheel build were also exercised in the preceding packaging update, with all package files verified through wheel installation.
  • Repository pre-commit, vendored-source integrity and DCO checks passed. The GPU runtime uses PR Python sources and a CI native artifact with matching native/custom-op source trees, rather than a clean native build of this PR.
  • A compiled 480x832x165, 20-step Wan comparison passed bitwise decoded-video parity (LPIPS 0). Three warmed draws per case, with an initial and restored FA4-default control, found no meaningful end-to-end gain on the tested SM103 development device. Additional forced-1CTA/VANILLA pipeline runs were omitted at the headroom gate. Standalone FA4 CTA timings are recorded separately.
  • SM100, production B200/B300 hardware, distributed execution, clean native/container/Jenkins builds and package-index publication remain unvalidated.

Keep this draft opt-in as a validation prototype. Before productizing runtime search, compare it with a simple validated static FA4 policy: the existing SM100 noncausal h128 2CTA table uses frequency 10, so frequency 16 versus that default and 1CTA is a focused next experiment. Retain runtime tuning only if workload-dependent winner changes deliver meaningful end-to-end gains beyond the simpler policy. A global FA4 table change also needs coverage beyond VisualGen self-attention. Before an index-based release, publish the companion FA4 wheel to the configured NVIDIA package index and verify resolution there. This PR builds and archives it; no package has been published. The release channel must support its downstream local version (+trtllm.1).

Signed-off-by: Chang Liu <9713593+chang-l@users.noreply.github.com>
Signed-off-by: Chang Liu <9713593+chang-l@users.noreply.github.com>
@chang-l chang-l changed the title [None][feat] Demo VisualGen FA4 CTA and exp2 autotuning [None][feat] Add VisualGen FA4 autotuning with bundled patched FA4 Sep 28, 2026
Signed-off-by: Chang Liu <9713593+chang-l@users.noreply.github.com>
@chang-l chang-l changed the title [None][feat] Add VisualGen FA4 autotuning with bundled patched FA4 [None][feat] Add VisualGen FA4 autotuning and patched FA4 wheel build Sep 29, 2026
Signed-off-by: Chang Liu <9713593+chang-l@users.noreply.github.com>
Signed-off-by: Chang Liu <9713593+chang-l@users.noreply.github.com>
@chang-l chang-l closed this Sep 29, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant