[v2] document CCCL kernel_reduce.cuh SingleTile + scan GridThreshold=500 insights
From reading cccl_upstream/cub/cub/device/dispatch/kernels/kernel_reduce.cuh:
- SingleTile path: when num_partitions fits in one tile (always true for
BI-V100 attention with <=200 partitions), reduce uses single CTA.
Phase 2 is never the bottleneck.
- StableReductionOrder=false uses atomicAdd to skip pass 2 entirely.
Not applicable to attention (compound accumulator), but confirms
the Phase 2 architecture is correct.
From reading cccl_upstream/cub/cub/agent/single_pass_scan_operators.cuh:
- delay<Delay, GridThreshold=500>(): when gridDim.x < 500, ALL delay
strategies collapse to __threadfence_block(). BI-V100 scan grids
have at most ~12 blocks (100K/8448). ALL delay tuning (ns, dcid, l2w)
is irrelevant — bench_bi100.py's no_delay 'win' was actually noise
between identical __threadfence_block() calls.
From reading cccl_upstream/thrust/examples/summed_area_table.cu:
- inclusive_scan_by_key pattern for per-row operations maps to
per-sequence softmax denominator computation in paged_attention.
This commit is contained in:
@@ -211,10 +211,21 @@ def paged_attention_v2_pytorch(
|
||||
# =============================================================
|
||||
# Phase 2: Cross-partition reduction (CCCL binary_op pattern)
|
||||
#
|
||||
# This is the summary_statistics.binary_op pattern:
|
||||
# Combine (max_a, sum_a, out_a) ⊕ (max_b, sum_b, out_b)
|
||||
# using numerically stable log-sum-exp rescaling.
|
||||
# CCCL kernel_reduce.cuh insight: when grid_size fits in a single
|
||||
# tile (num_partitions <= threads * items_per_thread), the reduce
|
||||
# uses SingleTile path — one CTA, no temp buffer, no pass 2 kernel.
|
||||
#
|
||||
# For BI-V100 with 98 partitions (100K tokens / 1024 partition_size):
|
||||
# SingleTile threshold = 512 * 24 = 12288 >> 98 → always SingleTile
|
||||
# This means Phase 2 is never the bottleneck.
|
||||
#
|
||||
# CCCL single_pass_scan_operators.cuh insight: delay() has a
|
||||
# GridThreshold=500 gate. BI-V100 scan grids are always < 500 blocks,
|
||||
# so ALL delay strategies (no_delay, fixed_delay, exponential_backon)
|
||||
# collapse to __threadfence_block(). Delay tuning is irrelevant here.
|
||||
#
|
||||
# Phase 2 follows summary_statistics.cu binary_op: combine
|
||||
# (max_a, sum_a, out_a) ⊕ (max_b, sum_b, out_b) via log-sum-exp.
|
||||
# Fully vectorized — no loop over partitions.
|
||||
# =============================================================
|
||||
pm = max_logits[seq_idx, :, :num_partitions] # [H, P]
|
||||
|
||||
Reference in New Issue
Block a user