[docs] GridEvenShare work distribution — BI-V100 attention reduce needs only 9 CTAs

From reading cccl_upstream/cub/cub/grid/grid_even_share.cuh:

Key finding: For Qwen3.6 attention score reduction (100K seq_len),
with tile_items=12288 (512 threads × 24 items), only 9 CTAs are
needed. All fit in one wave on 16 SMs.

This means reduce tuning (items/threads) matters less than the V1/V2
dispatch choice in paged_attn.py. V1 uses a single CTA iterating
sequentially over all KV blocks, completely bypassing GridEvenShare's
parallel distribution. V2 would enable partition-based parallelism.

Also documents: RAKE (scan) vs STRIP_MINE (reduce) strategies,
'big shares' load balancing, and the SingleTile fast path for
short sequences.
This commit is contained in:
project_6
2026-08-05 03:36:44 +00:00
parent ce42a8579d
commit 5ca49d0e7c

View File

@@ -352,3 +352,61 @@ Benchmark result (bif=8, alg=1, pref=2, tpb=256, unrl=1, vsp2=1):
tpb=256 = matches CCCL default
vsp2=1 vec_size parameter (powers of 2, so vsp2=1 means vec_size=2)
speedups: 1.203199 1.058919 1.019168 (fp16, 1M/16M/64M)
---
## GridEvenShare Work Distribution (grid_even_share.cuh)
> Added: 2026-08-04
### Two strategies: RAKE vs STRIP_MINE
**RAKE** (scan uses this): consecutive tiles per block
block k gets tiles [k*avg .. k*avg + avg-1]
block_stride = TILE_ITEMS (contiguous, no gaps)
**STRIP_MINE** (reduce uses this): interleaved tiles
block k gets tiles k, k+grid_size, k+2*grid_size, ...
block_stride = grid_size * TILE_ITEMS (strided)
### Concrete numbers for BI-V100 attention score reduce
tile_items = 512 * 24 = 12288 (bi100_plus_float32_o4)
max_grid = 2 * 16 * 5 = 160 (occupancy * SMs * subscription)
| seq_len | total_tiles | grid_size | tiles/block | waves |
|---------|-------------|-----------|-------------|-------|
| 1K | 1 | 1 | 1 | 1 |
| 8K | 1 | 1 | 1 | 1 |
| 32K | 3 | 3 | 1 | 1 |
| 100K | 9 | 9 | 1 | 1 |
| 1M | 82 | 82 | 1 | 3 |
Even at 100K tokens, only 9 CTAs are needed everything fits in one
wave on 16 SMs. This means:
1. Reduce tuning (items/threads) matters less than expected there
are so few tiles that the per-tile overhead dominates, not throughput.
2. The V1/V2 choice in paged_attn.py matters MORE V1 doesn't use
GridEvenShare at all, it's a single CTA iterating sequentially.
V2's partition-based approach enables parallel reduction.
3. For short sequences (≤8K, 1 tile), SingleTile path triggers:
just 1 CTA, 1 kernel launch, no temp storage.
### The "big shares" distribution
GridEvenShare handles uneven tile counts:
avg_tiles_per_block = total_tiles / grid_size
big_shares = total_tiles % grid_size (blocks that get +1 tile)
For 100K tokens with 9 tiles and 9 blocks: avg=1, big_shares=0.
All blocks equal. No imbalance.
For 1M tokens with 82 tiles and 82 blocks: avg=1, big_shares=0.
Still perfectly balanced at 1 tile/block.
Only when max_grid_size limits grid_size do we get imbalance:
e.g., 200 tiles with max_grid=160 avg=1, big_shares=40 (40 blocks
get 2 tiles, 120 blocks get 1 tile).