[ENGINE] CCCL kernel_scan.cuh dual-algorithm dispatch: v1 (online norm) vs v2 (deferred norm)

Source: cccl_upstream/cub/cub/device/dispatch/kernels/kernel_scan.cuh
Target: vllm/attention/ops/prefix_prefill.py

CCCL kernel_scan.cuh implements compile-time algorithm selection:
  - lookback: AgentScan with delay_constructor_t (safe default)
  - lookahead: warpspeed pipeline stages (SM90+, deferred reduction)

Applied to prefix_prefill Triton kernels:
  - _fwd_kernel (v1) = lookback: online softmax norm per block
  - _fwd_kernel_flash_attn_v2 = lookahead: deferred normalization
    Saves (ctx_len / BLOCK_N) divisions per query row.

Before: v2 kernel NEVER called — dead code since initial commit.
After: v2 dispatched for standard Qwen3.6 path (no alibi, no sliding
window, power-of-2 head_dim, no FP8).

BI-V100: v2 saves 64 fdiv/row at ctx_len=4096, BLOCK_N=64.
This commit is contained in:
dylanyunlon
2026-08-07 06:36:55 +00:00
parent 35e85dbc67
commit 8002900af0

View File

@@ -828,6 +828,42 @@ if triton.__version__ >= "2.1.0":
if sliding_window is None or sliding_window <= 0:
sliding_window = 0
# ═══════════════════════════════════════════════════════════════
# CCCL kernel_scan.cuh dual-algorithm dispatch pattern
#
# kernel_scan.cuh line 110:
# if constexpr (active_policy.algorithm == ScanAlgorithm::lookahead)
# → device_scan_lookahead_body(...) // deferred reduction
# else
# → AgentScan(...).ConsumeRange(...) // online reduction
#
# In prefix_prefill, the same pattern maps to:
# _fwd_kernel = "lookback" path (online normalization per block)
# _fwd_kernel_flash_attn_v2 = "lookahead" path (deferred norm at end)
#
# v2 does acc_scale = alpha (no division) inside the loop, then
# acc = acc / l_i[:, None] once at the end. This saves
# (ctx_len / BLOCK_N) divisions per query row.
#
# For BI-V100 (16 SMs, limited IPC): fewer instructions per
# iteration = better pipeline utilization.
#
# Selection criteria (from CCCL):
# lookahead requires: SM90+, contiguous iterators, CUDA 12.8+
# lookback: always safe
#
# Our criteria:
# v2: no alibi, no sliding_window, no FP8, head_dim is power-of-2
# (no padding needed → avoids dim_mask overhead)
# v1: alibi, sliding_window, FP8, or non-power-of-2 head_dim
# ═══════════════════════════════════════════════════════════════
use_v2_kernel = (
alibi_slopes is None
and sliding_window == 0
and Lk == Lk_padded # head_dim is power of 2
and "fp8" not in kv_cache_dtype
)
if alibi_slopes is not None:
_fwd_kernel_alibi[grid](
q,
@@ -882,54 +918,110 @@ if triton.__version__ >= "2.1.0":
)
return
_fwd_kernel[grid](
q,
k,
v,
k_cache,
v_cache,
b_loc,
sm_scale,
k_scale,
v_scale,
b_start_loc,
b_seq_len,
b_ctx_len,
v_cache.shape[3],
k_cache.shape[4],
o,
b_loc.stride(0),
b_loc.stride(1),
q.stride(0),
q.stride(1),
q.stride(2),
k.stride(0),
k.stride(1),
k.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
o.stride(0),
o.stride(1),
o.stride(2),
k_cache.stride(0),
k_cache.stride(1),
k_cache.stride(2),
k_cache.stride(3),
k_cache.stride(
4), #[num_blocks, num_kv_heads, head_size/x, block_size, x]
v_cache.stride(0),
v_cache.stride(1),
v_cache.stride(2),
v_cache.stride(
3), #[num_blocks, num_kv_heads, head_size, block_size]
num_queries_per_kv=num_queries_per_kv,
BLOCK_M=BLOCK,
BLOCK_DMODEL=Lk,
BLOCK_DMODEL_PADDED=Lk_padded,
BLOCK_N=BLOCK,
SLIDING_WINDOW=sliding_window,
num_warps=NUM_WARPS,
num_stages=1,
)
if use_v2_kernel:
# CCCL "lookahead" path: deferred normalization
# _fwd_kernel_flash_attn_v2 accumulates unnormalized weights,
# then divides once at the end (acc / l_i). Fewer divisions
# per iteration = better ALU utilization on BI-V100's 16 SMs.
#
# CCCL kernel_scan.cuh parallel:
# device_scan_lookahead_body does batched prefix sums
# with pipeline stages, deferring partial sums.
# Our v2 does the same conceptually: defer softmax norm.
_fwd_kernel_flash_attn_v2[grid](
q,
k,
v,
k_cache,
v_cache,
b_loc,
sm_scale,
b_start_loc,
b_seq_len,
b_ctx_len,
v_cache.shape[3],
k_cache.shape[4],
o,
b_loc.stride(0),
b_loc.stride(1),
q.stride(0),
q.stride(1),
q.stride(2),
k.stride(0),
k.stride(1),
k.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
o.stride(0),
o.stride(1),
o.stride(2),
k_cache.stride(0),
k_cache.stride(1),
k_cache.stride(2),
k_cache.stride(3),
k_cache.stride(4),
v_cache.stride(0),
v_cache.stride(1),
v_cache.stride(2),
v_cache.stride(3),
num_queries_per_kv=num_queries_per_kv,
BLOCK_M=BLOCK,
BLOCK_DMODEL=Lk,
BLOCK_N=BLOCK,
num_warps=NUM_WARPS,
num_stages=1,
)
else:
# CCCL "lookback" path: online normalization (safe default)
# Handles: alibi, sliding window, FP8, non-power-of-2 head_dim
_fwd_kernel[grid](
q,
k,
v,
k_cache,
v_cache,
b_loc,
sm_scale,
k_scale,
v_scale,
b_start_loc,
b_seq_len,
b_ctx_len,
v_cache.shape[3],
k_cache.shape[4],
o,
b_loc.stride(0),
b_loc.stride(1),
q.stride(0),
q.stride(1),
q.stride(2),
k.stride(0),
k.stride(1),
k.stride(2),
v.stride(0),
v.stride(1),
v.stride(2),
o.stride(0),
o.stride(1),
o.stride(2),
k_cache.stride(0),
k_cache.stride(1),
k_cache.stride(2),
k_cache.stride(3),
k_cache.stride(
4),
v_cache.stride(0),
v_cache.stride(1),
v_cache.stride(2),
v_cache.stride(3),
num_queries_per_kv=num_queries_per_kv,
BLOCK_M=BLOCK,
BLOCK_DMODEL=Lk,
BLOCK_DMODEL_PADDED=Lk_padded,
BLOCK_N=BLOCK,
SLIDING_WINDOW=sliding_window,
num_warps=NUM_WARPS,
num_stages=1,
)
return