Skip to content

Feat/block extend review fixes - #6

Open
fdz-1999 wants to merge 37 commits into
mainfrom
feat/block-extend-review-fixes
Open

fdz-1999 wants to merge 37 commits into
mainfrom
feat/block-extend-review-fixes

Conversation

@fdz-1999

Copy link
Copy Markdown
Owner

📌 Description

🔍 Related Issues

🚀 Pull Request Checklist

Thank you for contributing to FlashInfer! Before we review your pull request, please make sure the following items are complete.

✅ Pre-commit Checks

  • I have installed pre-commit by running pip install pre-commit (or used your preferred method).
  • I have installed the hooks with pre-commit install.
  • I have run the hooks manually with pre-commit run --all-files and fixed any reported issues.

If you are unsure about how to set up pre-commit, see the pre-commit documentation.

🧪 Tests

  • Tests have been added or updated as needed.
  • All tests are passing (unittest, etc.).

Reviewer Notes

yzh119 and others added 30 commits November 20, 2025 00:57
…lease CI

Merge branch fa2-fa3-opt of git@code.alipay.com:deep-xpu/flashinfer.git into main
https://code.alipay.com/deep-xpu/flashinfer/pull_requests/6

Reviewed-by: 明泓 <mingliang.gml@antgroup.com>


* feat(dllm,ci): add Block Expanding Attention & PyPI release CI
* build: add date suffix to ant-deepxpu-flashinfer-python version (0.5.3.20260202)
* fix(jit): compare base version only to allow date/cuda suffix
* drop clone api
* build flashinfer-jit-cache
Add DLLM Block Extend Attention feature with tile-level skip
optimization
using native MaskMode::kBlockExpanding.

Core API:
- Single-request: block_extend_attention_with_offset() with q/kv offset
  support
- Batch: BatchBlockExtendRaggedOffsetWrapper,
  BatchBlockExtendPagedOffsetWrapper
- Cascade: 3-stage attention (current chunk + prefix + merge state)
- Support both JIT and AOT compilation

Tests:
- Precision: block extend vs custom_mask reference, cascade vs blockwise
  correctness
- Performance: FlashInfer Block Extend vs PyTorch Flex Attention
  benchmark
  - Context length sweep (1K-32K), block size alignment analysis
  - Significant speedup over Flex Attention with lower memory usage
- Validate dllm_block_size > 0 to reject zero and negative values
- Raise ValueError on unsupported dtype instead of silent fallback to fp16
- Preserve user's preferred backend across wrapper re-creation
- Track idtype to correctly invalidate plan when index dtype changes
- Defer backend auto-selection to wrappers instead of pre-resolving in cascade
- Warn when q_offsets is None but prefix exists in cascade attention
- Pass device to FA3 SM90 check and include device in module cache key
- Remove unused logits_soft_cap parameter from sglang_style_cascade_attention
- Fix causal=True comment to causal=False in sglang_style_cascade_attention
- Fix docstring function names (block_expanding_* -> block_extend_*)
- Add assert statements to test correctness checks
- Rename benchmark functions from test_* to bench_* to avoid pytest collection
- Fix missing trailing newline in .cuh and .jinja files
- Add test_dllm_blockwise_mask_attention.py (renamed from test_dllm_cascade_vs_blockwise_extend_attention.py)
- Add scripts/task_jit_run_tests_dllm.sh for CI
- Add DLLM Block Extend test jobs (A10G + H100) in pr-test.yml
- Remove test_dllm_cascade_vs_blockwise_extend_attention.py
- Remove test_dllm_vs_flex_attention.py
- Raise RuntimeError when FA3 is explicitly requested on non-SM90 GPUs
- Apply same guard to both select_best_backend and select_best_backend_paged
- Sort __all__ in dllm/__init__.py alphabetically (Ruff RUF022)

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
Assert that k_prefix and v_prefix are both provided or both None,
instead of silently ignoring a partially specified prefix.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
Resolved conflicts in 4 files:
- flashinfer/__init__.py: kept both main's new imports + feature's dllm import
- flashinfer/jit/attention/modules.py: kept both main's _fa2_head_dim_nvcc_flags + feature's mask_modes/has_q_block_expanding_offset
- include/flashinfer/utils.cuh: kept both main's additions + feature's kBlockExpanding dispatch + DEFINE_HAS_MEMBER traits
- include/flashinfer/attention/prefill.cuh: used main's refactored loop structure (VO split, KV shared SMEM, FP8 repack) as base, injected feature's kBlockExpanding logic into num_iterations/mask_iteration/logits_mask in all 3 kernel variants (SinglePrefill, RaggedPrefill, PagedPrefill)
Main branch introduced a major refactor to prefill.cuh: VO split (head dim
tiled computation) and KV shared SMEM (K/V time-sharing the same shared
memory buffer), which restructured the main loop of all three kernels
(SinglePrefill, RaggedKV, PagedKV). The feature branch's kBlockExpanding
logic was injected into num_iterations / mask_iteration / logits_mask,
conflicting with main's restructured loop in the following areas:

1. RaggedKV kernel loop body: main's loop restructuring moved
   update_mdo_states to a specific timing point before block.sync().
   During merge, this call was accidentally dropped. Restored it after
   logits_mask and before block.sync(), aligned with SinglePrefill/PagedKV.

2. PagedKV mask_iteration call: the 5th argument was mistakenly written as
   CTA_TILE_Q during merge. Changed to CTA_TILE_KV (this parameter divides
   KV token count by KV tile size to get tile count), consistent with the
   other two kernels.

3. needs_mask condition: after merge, kBlockExpanding branch was mixed with
   window_iteration via short-circuit ||, making semantics ambiguous.
   Refactored to explicit if-constexpr branches so kBlockExpanding only
   depends on mask_iteration, unaffected by window_iteration.

Also cleaned up dead code in block_expanding_prefill.cuh, fixed
CUTLASS_DEVICE indentation in mainloop.cuh, and added trailing newline
to sparse_mainloop.cuh.
…lation

Reviewer-flagged correctness (§3 of the design doc):

- §3.1 zero-visible-KV TMA deadlock (Hopper/FA3): the consumer's `store_zero`
  fast-path skipped `mma_f16`, so it never arrived at `shared_storage.barrier_O`,
  and the next non-zero tile's producer waited forever at
  `barrier_O.wait((work_idx+1)%2)` (mainloop.cuh). Added the same
  `barrier_O.arrive` (gated `work_idx != 0`, last warp, elect_one_sync) to the
  consumer zero-visible block, mirroring mainloop_mma.cuh:82-90. Producer and
  consumer already advance work_idx in lockstep (both skip the ++ for a zero
  tile), so parity stays consistent. FA2 is immune.
  Regression test `test_zero_visible_kv_no_hang` added (single-path
  all-invisible/partial/baseline + batch zero-then-normal) with a kv_offset-aware
  reference helper; requires SM90 to exercise the hang path.

- §3.2 idtype-in-URI: batch dLLM URIs previously embedded only (head_dim, dtype),
  so int32/int64 index builds of the same shape aliased the same generated source
  directory. `_get_batch_be_module_uri(head_dim, dtype, idtype)` now embeds
  `idx{i32|i64}` matching the standard `dtype_idx_{...}` scheme at
  modules.py:390, and validates idtype. Single-path URIs unchanged (offsets are
  scalars; no idtype axis).

§1 isolation (closed product):

- batch dLLM `mask_modes` fixed [0,1,2,3,4] -> [MaskMode.BLOCK_EXPANDING.value]
  at both call sites, so the dedicated dLLM batch URI compiles only mask-mode 4
  (single path already did this).
- closed product enforced: dtype {fp16,bf16}, idtype {int32,int64},
  head_dim {64,128} (gated in the URI builders); backend stays auto/fa2/fa3.
  fp8 is already a hard reject (not a coercion) on this branch.

Note: the reviewer's §1 "shared-config inflation" concern is largely a
misframing on this branch — aot.py is untouched, the default mask list stays
[0,1,2,3], and the jinja additions are inert guarded getters, not a new
mask-mode axis value. The real waste was the dedicated *batch* dLLM URI
enumerating five modes; that's fixed. Full design + status in
docs/block-extend-design-response.md.

Not yet done (follow-up tasks tracked in the design doc): standalone dLLM gen
function (remove `mask_modes` from the shared gen_customize_* /
get_customize_batch_prefill_module), and the §2 API convergence (fold the four
flashinfer/dllm/ APIs into a `block_diffusion=` option on the existing prefill
APIs while keeping a thin shim). Deferred per user direction; needs GPU
verification on SM90 before the API-surface refactor.

Co-Authored-By: Claude <noreply@anthropic.com>
…path API convergence

Reviewer design #1 (isolation), single-path half of §1.2 step 3:

- Add `gen_customize_block_extend_single_prefill_module` and the batch twin
  `gen_customize_block_extend_batch_prefill_module` in jit/attention/modules.py:
  dedicated dLLM front-ends that compile ONLY `MaskMode::kBlockExpanding` over
  the closed product (fp16/bf16 x head_dim 64/128), enforce that product up
  front, and delegate to the shared gen_customize_* with mask_modes fixed to
  [kBlockExpanding]. Exported from flashinfer.jit.attention.
- Route the single-path dLLM front-end (block_extend.py) through the dedicated
  gen instead of the shared `gen_customize_single_prefill_module(... mask_modes=[4])`.
  This is the "standalone config/variant class" the reviewers asked for: the new
  mask mode is a small dedicated entry point, not a value multiplied into the big
  shared prefill cartesian product. The shared gen_customize_* keeps its default
  mask list [0,1,2,3], so no existing prefill URI compiles mode 4 on either path.

Reviewer design #2 (API convergence), single-path §2.2.2:

- Delete the single-path's own module cache (_MODULE_CACHE_WITH_OFFSET), the
  _get_aot_path / _check_aot_available helpers, and the hand-rolled AOT-vs-JIT
  branch (which lacked JitSpec.build_and_load's file-lock guard). The single-path
  module build now delegates to JitSpec.build_and_load() exactly like every other
  single-prefill call. Drops the now-unused tvm_ffi/jit_env/Path imports.

Deferred (need GPU verification, not possible from this host): the batch half of
§1.2 step 3 (route BatchBlockExtend{Ragged,Paged}OffsetWrapper through the
dedicated batch gen fn) and task #7 / §2.2.1 (expose block_diffusion= as a plan
option on the existing BatchPrefill...Wrapper). Both thread the variant jit_args
build and offset tensors through the hottest BatchPrefill runtime dispatch path
(prefill.py:1657 / :2788) and cannot be GPU-verified here — full design and
status in docs/block-extend-design-response.md §5.

Python files ast.parse-clean. No GPU verification done.

Co-Authored-By: Claude <noreply@anthropic.com>
…iewer §2)

Reviewer design #1 (isolation), batch half of §1.2 step 3:
- Delegation switch in get_customize_batch_prefill_module: when mask_modes is
  fixed to [kBlockExpanding], route to gen_customize_block_extend_batch_prefill_module
  (dedicated front-end — compiles only mode 4 over the closed dLLM product),
  instead of the shared gen_customize_batch_prefill_module. The shared default
  mask list stays [0,1,2,3], so no existing prefill URI compiles mode 4. Exported
  the dedicated gens from flashinfer.jit.

Reviewer design #2 (API convergence), §2.2.1 batch:
- Added a named `block_diffusion=` mask option to BOTH
  BatchPrefillWithPagedKVCacheWrapper and BatchPrefillWithRaggedKVCacheWrapper:
  __init__(block_diffusion=, dllm_block_size=), plan(q_offsets=, kv_offsets=)
  with auto-flip to kBlockExpanding (mirroring prefix_len_ptr→MULTIITEMSCORING),
  and run() injects the offsets + sm_scale + dllm_block_size from self (via
  prepare_jit_additional_args). This is the reviewers' preferred shape — a mask
  option on the existing prefill API rather than a new API family users discover
  separately. The variant jit module is built lazily at plan() time (dtype/
  head_dim/idtype are only known then), sharing a new build_block_diffusion_jit_args
  helper in flashinfer/dllm/batch_block_extend.py.
- All edits are guarded by self._block_diffusion (default False), so non-dLLM
  BatchPrefill users are completely unaffected.
- The dedicated BatchBlockExtend{Ragged,Paged}OffsetWrapper classes are now thin
  shims that share build_block_diffusion_jit_args with the named option (no own
  module cache / no own AOT path — already gone; now also shared variant wiring).

Test: test_block_diffusion_named_option exercises block_diffusion=True on both
existing wrappers and cross-checks vs the FA2 reference.

Deferred (task #9, reviewer §2.3): the dedicated gpu-tests-dllm-a10g/h100 CI
lanes are KEPT for now (they run the dLLM test suite that covers this blind
block-diffusion code). Full lane deletion + folding the test into the standard
prefill test part needs multi-spot YAML needs-graph surgery best done with the
maintainers' test-suite wiring — see docs/block-extend-design-response.md §5.

Python ast.parse-clean. No GPU verification done (cannot compile/import
flashinfer or build CUDA on this host) — to be verified on a GPU server.

Co-Authored-By: Claude <noreply@anthropic.com>
…e (reviewer §2 complete)

Completes the reviewers' design #2 ("complete convergence") for the
single-request path: block-diffusion is now a mask option on the existing
single-prefill API, not a flashinfer/dllm/ API family users discover separately.

- single_prefill_with_kv_cache gains block_diffusion=, dllm_block_size=,
  q_offset=, kv_offset= (impl + both @overload stubs). When block_diffusion=True
  it builds the block-extend variant via get_block_extend_module_with_offset
  (dedicated gen, fixed mask_mode=kBlockExpanding, closed dLLM product) and runs
  through single_prefill_with_kv_cache_with_jit_module — identical to the
  previous dedicated function but exposed as a kwarg on the native API.
  Validates: dllm_block_size power-of-2, no custom_mask, no sliding window.
  Normal (block_diffusion=False) path untouched.
- flashinfer/dllm.block_extend_attention_with_offset is now a thin shim that
  delegates to single_prefill_with_kv_cache(..., block_diffusion=True, ...).
  Removed now-unused imports (single_prefill_with_kv_cache_with_jit_module,
  MaskMode). The variant-module builder (get_block_extend_module_with_offset)
  stays in flashinfer/dllm/ to keep the C++ variant decl strings out of
  prefill.py; prefill lazy-imports it only in the block_diffusion branch.

Test: test_block_diffusion_single_native_option exercises the native path on
three single-request configs (baseline, incremental-chunk, cascade-current with
high kv_offset → zero-visible first tile) and cross-checks vs the kv_offset-aware
reference AND vs the dedicated shim (must be identical).

With this, reviewer design #2 is fully converged: both batch wrappers AND
single_prefill_with_kv_cache expose block_diffusion as a native mask option;
flashinfer/dllm/ is reduced to thin convenience shims (no own module cache, no
own AOT/JIT branch, no separate API family semantics).

Python ast.parse-clean. No GPU verification done on this host (cannot
compile/import flashinfer) — to be verified on a GPU server.

Co-Authored-By: Claude <noreply@anthropic.com>
…rs + shape-change rebuild

Closes the two functional gaps between the native `block_diffusion=` path on
BatchPrefillWith{Paged,Ragged}KVCacheWrapper and the dedicated
BatchBlockExtend*OffsetWrapper shim (functional completeness / correctness):

(a) CUDA-graph offset buffers: __init__ now accepts q_offsets_buf / kv_offsets_buf
    (block_diffusion only); plan() copies q_offsets/kv_offsets into them under
    use_cuda_graph (parity with the shim, required for SGLang-dLLM cuda-graph
    capture where the graph must read stable buffer addresses).

(b) Shape-change rebuild: plan() now rebuilds the variant jit module when
    (head_dim, dtype, idtype) changes across plans (tracked via _bd_built_key).
    Previously a second plan() with a different shape silently reused the stale
    module — a correctness bug. Parity with the shim's rebuild-on-shape-change.

Both guarded by self._block_diffusion (default False → normal users untouched).
Remaining minor difference vs the shim: per-run sm_scale override (the shim
exposed it as a knob but no caller — incl. cascade / SGLang-dLLM — overrides
sm_scale per run; sm_scale is plan-time everywhere). Documented in the design doc.

Python ast.parse-clean. No GPU verification on this host — to be verified on a
GPU server.

Co-Authored-By: Claude <noreply@anthropic.com>
…t_prefill_params.cuh

Literal execution of the reviewers' §1 ask — "reverting the shared jinja/config
edits and moving the dispatch into a standalone config/variant class."

Removed the `get_q/kv_block_expanding_offset` convenience getters from:
- all 4 shared jinja config files (batch_prefill_customize_config.jinja,
  batch_prefill_sm90_customize_config.jinja, single_prefill_customize_config.jinja,
  single_prefill_sm90_customize_config.jinja)
- the shared C++ header default_prefill_params.cuh (3 structs: SinglePrefillParams,
  BatchPrefillRaggedParams, BatchPrefillPagedParams) — including the PR-added
  members `dllm_block_size` / `q_block_expanding_offset` / `kv_block_expanding_offset`
  (single) and `maybe_q/kv_block_expanding_offset` (batch) and their constructor
  init entries. These were dead code: the shared default params struct is never
  instantiated with MASK_MODE=kBlockExpanding (the standard kernels use the
  default mask list [0,1,2,3]; only the dedicated dLLM variant instantiates
  kBlockExpanding, and it uses the CUSTOMIZED params struct).

The FA2 kernel (prefill.cuh) now reads the offset fields directly inside the
existing `if constexpr (MASK_MODE == kBlockExpanding)` branches — only compiled
for the dLLM customized variant, never for existing URIs:
- batch: `((params.maybe_q_block_expanding_offset != nullptr) ? params.maybe_q_block_expanding_offset[idx] : 0)` (nullptr-check preserved from the old getter)
- single: `params.q_block_expanding_offset` (scalar; old getter had no check)
The fields themselves remain via the generic `additional_params` mechanism
(additional_tensor_names / additional_scalar_names), which is not block-expanding-
specific. The Hopper/FA3 path already read `additional_params` via SFINAE traits
(no getter), unchanged.

Result: the shared jinja + shared C++ params header now contain ZERO
block-expanding-specific code. The dLLM dispatch is standalone (dedicated
gen_customize_block_extend_* + delegation switch; offset fields via generic
additional_params; variant via inline variant_decl). This is the literal
"revert shared jinja/config edits + standalone dispatch" the reviewer asked for.

Cannot compile-verify on this host (no CUDA toolchain) — the FA2 kernel field
reads and the jinja/C++ removals are mechanical and preserve the old getter
semantics (incl. nullptr-check), but MUST be confirmed by a CUDA build on the
GPU server.

Co-Authored-By: Claude <noreply@anthropic.com>
Completes the literal §1 revert — the public shared gen_customize_single_prefill_module
and gen_customize_batch_prefill_module no longer expose a `mask_modes` parameter
(it was the last "shared config edit" the reviewer asked to revert). Impl-factored:
the gen body moved to private _gen_customize_*_prefill_module_impl(mask_modes)
(unchanged), the public gen calls _impl with the fixed standard list [0,1,2,3],
and the dedicated gen_customize_block_extend_*_prefill_module calls _impl with
[kBlockExpanding]. get_customize_batch_prefill_module retains mask_modes ONLY as
a dispatcher to route [4]→dedicated gen (else→public shared gen, no mask_modes).

No existing caller (decode.py, pod.py, gen_single/batch_prefill_module wrappers,
prefill.py) passed mask_modes, so all are unaffected. The shared default mask
list stays [0,1,2,3] → existing prefill URIs never compile mode 4 (reviewer's
inflation concern, already absent, now also impossible via the public API).

JIT/AOT alignment preserved: the _impl body is unchanged, so the dLLM variant
compiles the identical kernel as before and JitSpec.build_and_load() still does
AOT-load-if-.so-exists else JIT (identical to the original blockwise-mask path's
actual behavior — FLASHINFER_FORCE_JIT was ineffective in the original too, since
build_and_load loads AOT when the .so exists).

Python ast.parse-clean. No CUDA build on this host — the impl-factoring is
mechanical (body unchanged under _impl) but verify on the GPU server.

Co-Authored-By: Claude <noreply@anthropic.com>
fdz-1999 and others added 7 commits July 19, 2026 01:53
After the §1 revert removed the single-prefill jinja getters, the
`has_q_block_expanding_offset` jinja var (set in the FA2 single gen branch of
gen_customize_single_prefill_module) is no longer referenced by any jinja
template. Drop the dead assignment. (The Hopper SFINAE traits
`has_q_block_expanding_offset_v<AdditionalParams>` are unrelated C++ template
vars and stay.) ast.parse-clean.

Co-Authored-By: Claude <noreply@anthropic.com>
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.

3 participants