[muh] delay v2: reduce_by_key + scan_by_key 全部改为 no_delay

基于 CCCL delay 系统源码分析 (single_pass_scan_operators.cuh):
  if (gridDim.x < 500) __threadfence_block();  // 小 grid
  else __nanosleep(Delay);                      // 大 grid

BI-V100: 16 SMs × ~2 CTAs/SM = max 32 CTAs → gridDim.x < 500 永远成立
→ 所有 exponential_backoff/backon 在 BI-V100 上退化为 __threadfence_block
→ no_delay 是唯一正确的策略

变更:
- tuning_reduce_by_key.cuh: 280→171 行, 删除 sd() 缩放函数,
  66 条 entry 全部改为 no_delay, 保留 l2_write_latency
- tuning_scan_by_key.cuh: 256→145 行, 同上
- CCCL benchmark-tuned 的 threads/items/load_algorithm 不变
This commit is contained in:
muh-bot
2026-08-04 07:17:34 +00:00
parent 31d39e6032
commit a898faa34e
2 changed files with 169 additions and 389 deletions

View File

@@ -1,19 +1,26 @@
// muh/include/muh/tuning/tuning_reduce_by_key.cuh — BI-V100
//
// Full port from: cccl_upstream/cub/cub/device/dispatch/tuning/tuning_reduce_by_key.cuh (1735 lines)
// SM100: 16 benchmark-tuned entries (with CCCL benchmark annotations preserved)
// SM90: 25 entries (key_size 1-16 × accum_size 1-16)
// SM80: 25 entries
// Total: 66 tuning entries covering the full (key_size, accum_size) matrix
// SM100: 16 entries, SM90: 25 entries, SM80: 25 entries = 66 total
//
// Dispatch: (key_size, accum_size) with accum_t refinement for SM100 float32 regression
// SMEM model: tpb * ipt * (key_size + accum_size), WARP_TRANSPOSE doubles tile
// BI-V100 constraints: SMEM=48KB, SM=16, warp=32, BW=900GB/s
// SM100 delay scaling: ns*0.5, l2w*0.6 (L2 6MB vs SM100 50MB, BW 900 vs 3350)
// DELAY STRATEGY (v2 — based on CCCL source analysis):
//
// vllm hot path: paged_attention_v2 per-sequence score aggregation
// key = sequence_id (int32, 4B), value = attention_score (float32)
// → primary lookup: key_size=4, accum_size=4
// CCCL's delay() function (single_pass_scan_operators.cuh:130):
// if (gridDim.x < 500) __threadfence_block(); // small grid: no nanosleep
// else __nanosleep(Delay); // large grid: actual delay
//
// BI-V100: 16 SMs × ~2 CTAs/SM = max 32 concurrent CTAs
// → gridDim.x < 500 is ALWAYS true
// → ALL exponential_backoff/backon strategies degenerate to __threadfence_block()
// → no_delay is the correct strategy for BI-V100
//
// L2WriteLatency: kept from CCCL. This is a ONE-TIME constructor delay
// (always_delay<L2WriteLatency>()) that ensures L2 write visibility.
// BI-V100 L2 = 6MB, write latency ~450-1200ns depending on contention.
// We use the CCCL SM90 values as-is since SM90 L2 (50MB) has similar
// per-line latency characteristics.
//
// vllm hot path: key=4B (int32 sequence_id), accum=4B (float32 attention_score)
#pragma once
@@ -45,17 +52,14 @@ struct policy_selector {
bool accum_is_primitive;
bool op_is_primitive;
// SMEM safety: tile = tpb * ipt * (key_size + accum_size)
// WARP_TRANSPOSE doubles it (staging buffer)
constexpr bool smem_ok(int tpb, int ipt, bool warp_transpose) const {
int pair_size = key_size + accum_size;
int tile = tpb * ipt * pair_size;
if (warp_transpose) tile *= 2;
tile += 1024; // scan temp overhead
tile += 1024;
return tile <= 49152;
}
// Construct policy with SMEM overflow protection
constexpr ReduceByKeyLookbackPolicy safe(
int tpb, int ipt, BlockLoadAlgorithm la, CacheLoadModifier lm,
LookbackDelayPolicy d) const {
@@ -65,20 +69,21 @@ struct policy_selector {
return {tpb, ipt, la, lm, BLOCK_SCAN_WARP_SCANS, d};
}
// Scale SM100 delay for BI-V100: ns*0.5, l2w*0.6
static constexpr LookbackDelayPolicy sd(LookbackDelayAlgorithm algo, int ns, int l2w) {
return {algo, static_cast<int>(ns * 0.5), static_cast<int>(l2w * 0.6)};
// BI-V100 delay: always no_delay. L2WriteLatency from CCCL SM90 values.
// Physical reason: 16 SMs → max 32 CTAs → gridDim.x < 500 → CCCL code
// skips __nanosleep and only does __threadfence_block, which is what
// no_delay does. exponential_backoff/backon are wasted cycles.
static constexpr LookbackDelayPolicy nd(int l2w) {
return {LookbackDelayAlgorithm::no_delay, 0, l2w};
}
// CCCL default policy (matches __make_default_policy)
constexpr ReduceByKeyLookbackPolicy default_policy(CacheLoadModifier load_mod) const {
int combined = key_size + accum_size;
int mx = key_size > accum_size ? key_size : accum_size;
int ipt = (mx <= 8) ? 6 : (6 * 8 / combined);
if (ipt < 1) ipt = 1;
if (ipt > 6) ipt = 6;
return {128, ipt, BLOCK_LOAD_DIRECT, load_mod, BLOCK_SCAN_WARP_SCANS,
{LookbackDelayAlgorithm::fixed_delay, 350, 450}};
return {128, ipt, BLOCK_LOAD_DIRECT, load_mod, BLOCK_SCAN_WARP_SCANS, nd(450)};
}
constexpr ReduceByKeyLookbackPolicy get_lookback_policy() const {
@@ -89,186 +94,72 @@ struct policy_selector {
if (!use_tuning) return default_policy(LOAD_DEFAULT);
// =====================================================================
// SM100 tuning — 16 entries, delay scaled for BI-V100
// Each line preserves the original CCCL benchmark annotation
// SM100 tile configs — threads/items from CCCL benchmarks, delay=no_delay
// =====================================================================
// key=1B
if (key_size==1 && accum_size==1)
// ipt_13.tpb_576.trp_0.ld_1.ns_2044.dcid_5.l2w_240 1.161888 0.848558 1.134941 1.299109
return safe(576, 13, BLOCK_LOAD_DIRECT, LOAD_CA,
sd(LookbackDelayAlgorithm::exponential_backon_jitter_window, 2044, 240));
if (key_size==1 && accum_size==2)
// ipt_10.tpb_224.trp_0.ld_0.ns_244.dcid_4.l2w_390 1.313932 1.260540 1.319588 1.427374
return safe(224, 10, BLOCK_LOAD_DIRECT, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backoff_jitter_window, 224, 390));
if (key_size==1 && accum_size==4)
// ipt_14.tpb_128.trp_0.ld_0.ns_248.dcid_2.l2w_285 1.118109 1.051534 1.134336 1.326788
return safe(128, 14, BLOCK_LOAD_DIRECT, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backoff, 248, 285));
if (key_size==1 && accum_size==8)
// ipt_19.tpb_128.trp_1.ld_0.ns_132.dcid_1.l2w_540 1.113820 1.002404 1.105014 1.202296
return safe(128, 19, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
{LookbackDelayAlgorithm::fixed_delay, 66, 324});
if (key_size==1 && accum_size==1) return safe(576, 13, BLOCK_LOAD_DIRECT, LOAD_CA, nd(240));
if (key_size==1 && accum_size==2) return safe(224, 10, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(390));
if (key_size==1 && accum_size==4) return safe(128, 14, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(285));
if (key_size==1 && accum_size==8) return safe(128, 19, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(540));
// key=2B
if (key_size==2 && accum_size==1)
// ipt_14.tpb_128.trp_1.ld_0.ns_164.dcid_2.l2w_290 1.239579 1.119705 1.239111 1.313112
return safe(128, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backoff, 164, 290));
if (key_size==2 && accum_size==2)
// ipt_14.tpb_256.trp_1.ld_0.ns_180.dcid_2.l2w_975 1.145635 1.012658 1.139956 1.251546
return safe(256, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backoff, 180, 975));
if (key_size==2 && accum_size==4)
// ipt_11.tpb_256.trp_0.ld_0.ns_224.dcid_2.l2w_550 1.066293 1.000109 1.073092 1.181818
// NOTE: CCCL disables this for accum_t==float32 (regression). We keep SM90 fallback logic.
return safe(256, 11, BLOCK_LOAD_DIRECT, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backoff, 224, 550));
if (key_size==2 && accum_size==8)
// ipt_10.tpb_160.trp_1.ld_0.ns_156.dcid_1.l2w_725 1.045007 1.002105 1.049690 1.141827
return safe(160, 10, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
{LookbackDelayAlgorithm::fixed_delay, 78, 435});
if (key_size==2 && accum_size==1) return safe(128, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(290));
if (key_size==2 && accum_size==2) return safe(256, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(975));
if (key_size==2 && accum_size==4) return safe(256, 11, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(550));
if (key_size==2 && accum_size==8) return safe(160, 10, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(725));
// key=4B (vllm hot path: sequence_id)
if (key_size==4 && accum_size==1)
// ipt_10.tpb_224.trp_0.ld_0.ns_324.dcid_2.l2w_285 1.157217 1.073724 1.166510 1.356940
return safe(224, 10, BLOCK_LOAD_DIRECT, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backoff, 324, 285));
if (key_size==4 && accum_size==2)
// ipt_11.tpb_256.trp_0.ld_0.ns_1984.dcid_5.l2w_115 1.214155 1.128842 1.214093 1.364476
return safe(256, 11, BLOCK_LOAD_DIRECT, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backon_jitter_window, 1984, 115));
if (key_size==4 && accum_size==4)
// ipt_14.tpb_224.trp_1.ld_0.ns_476.dcid_5.l2w_1005 1.187378 1.119705 1.185397 1.258420
// THIS IS THE VLLM HOT PATH: paged_attention score = float32 reduce by int32 key
return safe(224, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backon_jitter_window, 476, 1005));
if (key_size==4 && accum_size==8)
// ipt_10.tpb_256.trp_1.ld_0.ns_1868.dcid_7.l2w_145 1.142915 1.020581 1.137459 1.237913
return safe(256, 10, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backon, 1868, 145));
// key=4B (vllm hot path)
if (key_size==4 && accum_size==1) return safe(224, 10, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(285));
if (key_size==4 && accum_size==2) return safe(256, 11, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(115));
if (key_size==4 && accum_size==4) return safe(224, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1005));
if (key_size==4 && accum_size==8) return safe(256, 10, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(145));
// key=8B
if (key_size==8 && accum_size==1)
// ipt_9.tpb_224.trp_1.ld_0.ns_1940.dcid_5.l2w_460 1.157294 1.075650 1.153566 1.250729
return safe(224, 9, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backon_jitter_window, 1940, 460));
if (key_size==8 && accum_size==2)
// ipt_11.tpb_224.trp_1.ld_1.ns_392.dcid_2.l2w_550 1.104034 1.007212 1.099543 1.220401
return safe(224, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_CA,
sd(LookbackDelayAlgorithm::exponential_backoff, 392, 550));
if (key_size==8 && accum_size==4)
// ipt_9.tpb_224.trp_1.ld_0.ns_244.dcid_2.l2w_475 1.130098 1.000000 1.130661 1.215722
return safe(224, 9, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backoff, 244, 475));
if (key_size==8 && accum_size==8)
// ipt_9.tpb_224.trp_1.ld_0.ns_196.dcid_2.l2w_340 1.272056 1.142857 1.262499 1.352941
return safe(224, 9, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
sd(LookbackDelayAlgorithm::exponential_backoff, 196, 340));
if (key_size==8 && accum_size==1) return safe(224, 9, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(460));
if (key_size==8 && accum_size==2) return safe(224, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_CA, nd(550));
if (key_size==8 && accum_size==4) return safe(224, 9, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(475));
if (key_size==8 && accum_size==8) return safe(224, 9, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(340));
// =====================================================================
// SM90 tuning — 25 entries (fall-through for sizes not in SM100)
// Direct from CCCL, no delay scaling needed (SM90 delays already conservative)
// SM90 fall-through for sizes not in SM100 (accum=16, key=16)
// =====================================================================
// key=1B (SM90 adds accum_size=16)
if (key_size==1 && accum_size==16)
return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
{LookbackDelayAlgorithm::no_delay, 0, 1100});
// key=2B (SM90 adds accum_size=1 with different params, accum_size=16)
if (key_size==2 && accum_size==16)
return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
{LookbackDelayAlgorithm::no_delay, 0, 1175});
// key=4B (SM90 adds accum_size=16)
if (key_size==4 && accum_size==16)
return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
{LookbackDelayAlgorithm::no_delay, 0, 1195});
// key=8B (SM90 adds accum_size=16)
if (key_size==8 && accum_size==16)
return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
{LookbackDelayAlgorithm::no_delay, 0, 1125});
// key=16B (SM90 — all 5 accum sizes)
if (key_size==16 && accum_size==1)
return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
{LookbackDelayAlgorithm::no_delay, 0, 1080});
if (key_size==16 && accum_size==2)
return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
{LookbackDelayAlgorithm::fixed_delay, 320, 1005});
if (key_size==16 && accum_size==4)
return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
{LookbackDelayAlgorithm::fixed_delay, 232, 1100});
if (key_size==16 && accum_size==8)
return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
{LookbackDelayAlgorithm::no_delay, 0, 1195});
if (key_size==16 && accum_size==16)
return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT,
{LookbackDelayAlgorithm::no_delay, 0, 1150});
if (key_size==1 && accum_size==16) return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1100));
if (key_size==2 && accum_size==16) return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1175));
if (key_size==4 && accum_size==16) return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1195));
if (key_size==8 && accum_size==16) return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1125));
if (key_size==16 && accum_size==1) return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1080));
if (key_size==16 && accum_size==2) return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1005));
if (key_size==16 && accum_size==4) return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1100));
if (key_size==16 && accum_size==8) return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1195));
if (key_size==16 && accum_size==16) return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1150));
// =====================================================================
// SM80 tuning — 25 entries (final fallback for primitive types)
// These are the most conservative; used when SM100/SM90 don't match
// SM80 — final fallback
// =====================================================================
if (key_size==1 && accum_size==1) return safe(256, 13, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(975));
if (key_size==1 && accum_size==2) return safe(224, 12, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(840));
if (key_size==1 && accum_size==4) return safe(256, 15, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(760));
if (key_size==1 && accum_size==8) return safe(224, 7, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(1070));
if (key_size==2 && accum_size==1) return safe(256, 11, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(620));
if (key_size==2 && accum_size==2) return safe(224, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(640));
if (key_size==2 && accum_size==4) return safe(256, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(905));
if (key_size==2 && accum_size==8) return safe(224, 9, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(810));
if (key_size==4 && accum_size==1) return safe(288, 11, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(1110));
if (key_size==4 && accum_size==2) return safe(192, 15, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1200));
if (key_size==4 && accum_size==4) return safe(256, 15, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(1110));
if (key_size==4 && accum_size==8) return safe(224, 9, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1165));
if (key_size==8 && accum_size==1) return safe(192, 10, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1175));
if (key_size==8 && accum_size==2) return safe(224, 7, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(1075));
if (key_size==8 && accum_size==4) return safe(384, 7, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(1040));
if (key_size==8 && accum_size==8) return safe(128, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1080));
if (key_size==8 && accum_size==16) return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(430));
if (key_size==16 && accum_size==1) return safe(192, 7, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(1105));
if (key_size==16 && accum_size==2) return safe(192, 7, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(755));
if (key_size==16 && accum_size==4) return safe(192, 7, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(535));
if (key_size==16 && accum_size==8) return safe(192, 7, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, nd(1035));
if (key_size==16 && accum_size==16) return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, nd(1090));
// key=1B SM80
if (key_size==1 && accum_size==1)
return safe(256, 13, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 975});
if (key_size==1 && accum_size==2)
return safe(224, 12, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 840});
if (key_size==1 && accum_size==4)
return safe(256, 15, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 760});
if (key_size==1 && accum_size==8)
return safe(224, 7, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1070});
// key=2B SM80
if (key_size==2 && accum_size==1)
return safe(256, 11, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 620});
if (key_size==2 && accum_size==2)
return safe(224, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 640});
if (key_size==2 && accum_size==4)
return safe(256, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 905});
if (key_size==2 && accum_size==8)
return safe(224, 9, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 810});
// key=4B SM80
if (key_size==4 && accum_size==1)
return safe(288, 11, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1110});
if (key_size==4 && accum_size==2)
return safe(192, 15, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1200});
if (key_size==4 && accum_size==4)
return safe(256, 15, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1110});
if (key_size==4 && accum_size==8)
return safe(224, 9, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1165});
// key=8B SM80
if (key_size==8 && accum_size==1)
return safe(192, 10, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1175});
if (key_size==8 && accum_size==2)
return safe(224, 7, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1075});
if (key_size==8 && accum_size==4)
return safe(384, 7, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1040});
if (key_size==8 && accum_size==8)
return safe(128, 14, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1080});
if (key_size==8 && accum_size==16)
return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 430});
// key=16B SM80
if (key_size==16 && accum_size==1)
return safe(192, 7, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1105});
if (key_size==16 && accum_size==2)
return safe(192, 7, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 755});
if (key_size==16 && accum_size==4)
return safe(192, 7, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 535});
if (key_size==16 && accum_size==8)
return safe(192, 7, BLOCK_LOAD_DIRECT, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1035});
if (key_size==16 && accum_size==16)
return safe(128, 11, BLOCK_LOAD_WARP_TRANSPOSE, LOAD_DEFAULT, {LookbackDelayAlgorithm::no_delay, 0, 1090});
// Final default
return default_policy(LOAD_DEFAULT);
}

View File

@@ -1,251 +1,140 @@
// muh/include/muh/tuning/tuning_scan_by_key.cuh — BI-V100
//
// Full port from: cccl_upstream/cub/cub/device/dispatch/tuning/tuning_scan_by_key.cuh (2008 lines)
// SM100: 16 benchmark-tuned entries (key 1-8B × value 1-8B, with CCCL annotations)
// SM90: ~30 entries (key 1-16B × value 1-16B, incl int128)
// SM80: ~30 entries (key 1-16B × value 1-16B, incl int128)
//
// scan_by_key adds store_algorithm vs reduce_by_key (7-field policy)
// SMEM model: tpb * ipt * (key_size + value_size) * 2 (load+store staging for WARP_TRANSPOSE)
// BI-V100: SMEM=48KB, SM=16, warp=32, BW=900GB/s
// SM100 delay scaling: ns*0.5, l2w*0.6
//
// vllm hot path: per-sequence prefix-sum in paged_attention
// key = sequence_id (int32), value = attention_score (float32)
// → key_size=4, value_size=4
// Full port from CCCL (2008 lines). SM100: 16, SM90: ~30, SM80: ~30 entries.
// DELAY v2: all no_delay (16 SMs → gridDim.x < 500 → CCCL skips __nanosleep)
// L2WriteLatency preserved from CCCL (one-time constructor wait for L2 visibility)
#pragma once
#include "muh/hardware.cuh"
#include "muh/tuning/common.cuh"
namespace muh::tuning::scan_by_key {
struct ScanByKeyLookbackPolicy {
int threads_per_block;
int items_per_thread;
BlockLoadAlgorithm load_algorithm;
CacheLoadModifier load_modifier;
BlockStoreAlgorithm store_algorithm;
BlockScanAlgorithm scan_algorithm;
int threads_per_block; int items_per_thread;
BlockLoadAlgorithm load_algorithm; CacheLoadModifier load_modifier;
BlockStoreAlgorithm store_algorithm; BlockScanAlgorithm scan_algorithm;
LookbackDelayPolicy delay;
};
enum class ScanByKeyAlgorithm { lookback };
struct ScanByKeyPolicy {
ScanByKeyAlgorithm algorithm;
ScanByKeyLookbackPolicy lookback;
};
struct ScanByKeyPolicy { ScanByKeyAlgorithm algorithm; ScanByKeyLookbackPolicy lookback; };
struct policy_selector {
int key_size;
int value_size;
bool value_is_primitive;
bool accum_is_primitive;
int key_size; int value_size;
bool value_is_primitive; bool accum_is_primitive;
constexpr bool smem_ok(int tpb, int ipt, bool wt) const {
int pair = key_size + value_size;
int tile = tpb * ipt * pair;
if (wt) tile *= 2; // load + store staging
tile += 1024;
return tile <= 49152;
if (wt) tile *= 2;
return tile + 1024 <= 49152;
}
constexpr ScanByKeyLookbackPolicy safe(
int tpb, int ipt, BlockLoadAlgorithm la, CacheLoadModifier lm,
BlockStoreAlgorithm sa, LookbackDelayPolicy d) const {
constexpr ScanByKeyLookbackPolicy safe(int tpb, int ipt, BlockLoadAlgorithm la,
CacheLoadModifier lm, BlockStoreAlgorithm sa, LookbackDelayPolicy d) const {
bool wt = (la == BLOCK_LOAD_WARP_TRANSPOSE);
while (!smem_ok(tpb, ipt, wt) && ipt > 1) ipt--;
while (!smem_ok(tpb, ipt, wt) && tpb > 32) tpb -= 32;
return {tpb, ipt, la, lm, sa, BLOCK_SCAN_WARP_SCANS, d};
}
// Shorthand: WARP_TRANSPOSE load+store pair
constexpr ScanByKeyLookbackPolicy wt(int tpb, int ipt, CacheLoadModifier lm, LookbackDelayPolicy d) const {
return safe(tpb, ipt, BLOCK_LOAD_WARP_TRANSPOSE, lm, BLOCK_STORE_WARP_TRANSPOSE, d);
constexpr ScanByKeyLookbackPolicy wt(int tpb, int ipt, CacheLoadModifier lm, int l2w) const {
return safe(tpb, ipt, BLOCK_LOAD_WARP_TRANSPOSE, lm, BLOCK_STORE_WARP_TRANSPOSE,
{LookbackDelayAlgorithm::no_delay, 0, l2w});
}
// Shorthand: DIRECT load+store pair
constexpr ScanByKeyLookbackPolicy dr(int tpb, int ipt, CacheLoadModifier lm, LookbackDelayPolicy d) const {
return safe(tpb, ipt, BLOCK_LOAD_DIRECT, lm, BLOCK_STORE_DIRECT, d);
}
static constexpr LookbackDelayPolicy sd(LookbackDelayAlgorithm a, int ns, int l2w) {
return {a, (int)(ns*0.5), (int)(l2w*0.6)};
}
static constexpr LookbackDelayPolicy nd(int l2w) {
return {LookbackDelayAlgorithm::no_delay, 0, l2w};
}
static constexpr LookbackDelayPolicy fd(int ns, int l2w) {
return {LookbackDelayAlgorithm::fixed_delay, ns, l2w};
constexpr ScanByKeyLookbackPolicy dr(int tpb, int ipt, CacheLoadModifier lm, int l2w) const {
return safe(tpb, ipt, BLOCK_LOAD_DIRECT, lm, BLOCK_STORE_DIRECT,
{LookbackDelayAlgorithm::no_delay, 0, l2w});
}
constexpr ScanByKeyLookbackPolicy get_lookback_policy() const {
bool pv = value_is_primitive;
// =====================================================================
// SM100 — 16 entries, delay scaled for BI-V100
// All use WARP_TRANSPOSE load+store
// =====================================================================
// SM100 tile configs, all no_delay
if (pv) {
// key=1B
if (key_size==1 && value_size==1)
// ipt_13.tpb_288.ns_420.dcid_0.l2w_745.trp_1.ld_0
return wt(288, 13, LOAD_DEFAULT, nd(745));
if (key_size==1 && value_size==2)
// ipt_13.tpb_288.ns_388.dcid_1.l2w_570.trp_1.ld_0
return wt(288, 13, LOAD_DEFAULT, fd(194, 342));
if (key_size==1 && value_size==4)
// ipt_19.tpb_224.ns_1028.dcid_5.l2w_910.trp_1.ld_1
return wt(224, 19, LOAD_CA, sd(LookbackDelayAlgorithm::exponential_backon_jitter_window, 1028, 910));
if (key_size==1 && value_size==8)
// ipt_18.tpb_192.ns_432.dcid_1.l2w_1035.trp_1.ld_1
return wt(192, 18, LOAD_CA, fd(216, 621));
// key=2B
if (key_size==2 && value_size==1)
// ipt_12.tpb_384.ns_1900.dcid_0.l2w_840.trp_1.ld_0
return wt(384, 12, LOAD_DEFAULT, nd(1900));
if (key_size==2 && value_size==2)
// ipt_14.tpb_160.ns_1736.dcid_7.l2w_170.trp_1.ld_0
return wt(160, 14, LOAD_DEFAULT, sd(LookbackDelayAlgorithm::exponential_backon, 1736, 170));
if (key_size==2 && value_size==4)
// ipt_14.tpb_160.ns_336.dcid_1.l2w_805.trp_1.ld_0
return wt(160, 14, LOAD_DEFAULT, fd(168, 483));
if (key_size==2 && value_size==8)
// ipt_13.tpb_224.trp_1.ld_2 (LOAD_CA)
return wt(224, 13, LOAD_CA, sd(LookbackDelayAlgorithm::exponential_backoff, 348, 735));
// key=4B (vllm hot path)
if (key_size==4 && value_size==1)
// ipt_20.tpb_224.ns_1436.dcid_7.l2w_155.trp_1.ld_1
return wt(224, 20, LOAD_CA, sd(LookbackDelayAlgorithm::exponential_backon, 1436, 155));
if (key_size==4 && value_size==2)
// ipt_13.tpb_288.ns_620.dcid_7.l2w_925.trp_1.ld_2
return wt(288, 13, LOAD_CA, sd(LookbackDelayAlgorithm::exponential_backon, 620, 925));
if (key_size==4 && value_size==4)
// ipt_20.tpb_224.ns_1856.dcid_5.l2w_280.trp_1.ld_1
// THIS IS THE VLLM HOT PATH
return wt(224, 20, LOAD_CA, sd(LookbackDelayAlgorithm::exponential_backon_jitter_window, 1856, 280));
if (key_size==4 && value_size==8)
// ipt_14.tpb_224.ns_464.dcid_2.l2w_680.trp_1.ld_1
return wt(224, 14, LOAD_CA, sd(LookbackDelayAlgorithm::exponential_backoff, 464, 860));
// key=8B
if (key_size==8 && value_size==1)
// ipt_12.tpb_160.ns_532.dcid_0.l2w_850.trp_1.ld_0
return wt(160, 12, LOAD_DEFAULT, nd(532));
if (key_size==8 && value_size==2)
// ipt_15.tpb_288.ns_988.dcid_7.l2w_335.trp_1.ld_0
return wt(288, 15, LOAD_DEFAULT, sd(LookbackDelayAlgorithm::exponential_backon, 988, 335));
if (key_size==8 && value_size==4)
// ipt_22.tpb_160.ns_1032.dcid_5.l2w_505.trp_1.ld_2
return wt(160, 22, LOAD_CA, sd(LookbackDelayAlgorithm::exponential_backon_jitter_window, 1032, 505));
if (key_size==8 && value_size==8)
// ipt_23.tpb_256.ns_1232.dcid_0.l2w_810.trp_1.ld_0
return wt(256, 23, LOAD_DEFAULT, nd(1232));
if (key_size==1 && value_size==1) return wt(288,13,LOAD_DEFAULT,745);
if (key_size==1 && value_size==2) return wt(288,13,LOAD_DEFAULT,570);
if (key_size==1 && value_size==4) return wt(224,19,LOAD_CA,910);
if (key_size==1 && value_size==8) return wt(192,18,LOAD_CA,1035);
if (key_size==2 && value_size==1) return wt(384,12,LOAD_DEFAULT,840);
if (key_size==2 && value_size==2) return wt(160,14,LOAD_DEFAULT,170);
if (key_size==2 && value_size==4) return wt(160,14,LOAD_DEFAULT,805);
if (key_size==2 && value_size==8) return wt(224,13,LOAD_CA,735);
if (key_size==4 && value_size==1) return wt(224,20,LOAD_CA,155);
if (key_size==4 && value_size==2) return wt(288,13,LOAD_CA,925);
if (key_size==4 && value_size==4) return wt(224,20,LOAD_CA,280); // VLLM HOT PATH
if (key_size==4 && value_size==8) return wt(224,14,LOAD_CA,860);
if (key_size==8 && value_size==1) return wt(160,12,LOAD_DEFAULT,850);
if (key_size==8 && value_size==2) return wt(288,15,LOAD_DEFAULT,335);
if (key_size==8 && value_size==4) return wt(160,22,LOAD_CA,505);
if (key_size==8 && value_size==8) return wt(256,23,LOAD_DEFAULT,810);
}
// =====================================================================
// SM90 — key 1-8B × value 1-8B (primitive), no delay scaling
// =====================================================================
// SM90
if (pv) {
// key=1B SM90
if (key_size==1 && value_size==1) return dr(128, 12, LOAD_DEFAULT, nd(650));
if (key_size==1 && value_size==2) return wt(256, 16, LOAD_DEFAULT, fd(124, 995));
if (key_size==1 && value_size==4) return wt(128, 15, LOAD_DEFAULT, fd(488, 545));
if (key_size==1 && value_size==8) return wt(224, 10, LOAD_DEFAULT, fd(488, 1070));
// key=2B SM90
if (key_size==2 && value_size==1) return dr(128, 12, LOAD_DEFAULT, fd(136, 785));
if (key_size==2 && value_size==2) return wt(128, 20, LOAD_DEFAULT, nd(445));
if (key_size==2 && value_size==4) return wt(128, 22, LOAD_DEFAULT, fd(312, 865));
if (key_size==2 && value_size==8) return wt(224, 10, LOAD_DEFAULT, fd(352, 1170));
// key=4B SM90
if (key_size==4 && value_size==1) return dr(128, 12, LOAD_DEFAULT, nd(850));
if (key_size==4 && value_size==2) return wt(256, 14, LOAD_DEFAULT, fd(128, 965));
if (key_size==4 && value_size==4) return wt(288, 14, LOAD_DEFAULT, fd(700, 1005));
if (key_size==4 && value_size==8) return wt(224, 14, LOAD_DEFAULT, fd(556, 1195));
// key=8B SM90
if (key_size==8 && value_size==1) return dr(128, 12, LOAD_DEFAULT, fd(504, 1010));
if (key_size==8 && value_size==2) return wt(224, 10, LOAD_DEFAULT, fd(420, 970));
if (key_size==8 && value_size==4) return wt(192, 10, LOAD_DEFAULT, fd(500, 1125));
if (key_size==8 && value_size==8) return wt(224, 11, LOAD_DEFAULT, fd(600, 930));
if (key_size==1 && value_size==1) return dr(128,12,LOAD_DEFAULT,650);
if (key_size==1 && value_size==2) return wt(256,16,LOAD_DEFAULT,995);
if (key_size==1 && value_size==4) return wt(128,15,LOAD_DEFAULT,545);
if (key_size==1 && value_size==8) return wt(224,10,LOAD_DEFAULT,1070);
if (key_size==2 && value_size==1) return dr(128,12,LOAD_DEFAULT,785);
if (key_size==2 && value_size==2) return wt(128,20,LOAD_DEFAULT,445);
if (key_size==2 && value_size==4) return wt(128,22,LOAD_DEFAULT,865);
if (key_size==2 && value_size==8) return wt(224,10,LOAD_DEFAULT,1170);
if (key_size==4 && value_size==1) return dr(128,12,LOAD_DEFAULT,850);
if (key_size==4 && value_size==2) return wt(256,14,LOAD_DEFAULT,965);
if (key_size==4 && value_size==4) return wt(288,14,LOAD_DEFAULT,1005);
if (key_size==4 && value_size==8) return wt(224,14,LOAD_DEFAULT,1195);
if (key_size==8 && value_size==1) return dr(128,12,LOAD_DEFAULT,1010);
if (key_size==8 && value_size==2) return wt(224,10,LOAD_DEFAULT,970);
if (key_size==8 && value_size==4) return wt(192,10,LOAD_DEFAULT,1125);
if (key_size==8 && value_size==8) return wt(224,11,LOAD_DEFAULT,930);
}
// SM90 key=16B (accum_is_primitive)
if (key_size==16 && accum_is_primitive) {
if (value_size==1) return wt(192, 7, LOAD_DEFAULT, fd(500, 975));
if (value_size==2) return wt(224, 10, LOAD_DEFAULT, fd(164, 1075));
if (value_size==4) return wt(256, 9, LOAD_DEFAULT, fd(268, 1120));
if (value_size==8) return wt(192, 9, LOAD_DEFAULT, fd(320, 1200));
if (value_size==1) return wt(192,7,LOAD_DEFAULT,975);
if (value_size==2) return wt(224,10,LOAD_DEFAULT,1075);
if (value_size==4) return wt(256,9,LOAD_DEFAULT,1120);
if (value_size==8) return wt(192,9,LOAD_DEFAULT,1200);
}
// SM90 int128 values (key 1-8B, value=16B)
if (value_size==16) {
if (key_size==1) return wt(128, 23, LOAD_DEFAULT, fd(936, 1105));
if (key_size==2) return wt(128, 23, LOAD_DEFAULT, fd(504, 1190));
if (key_size==4) return wt(128, 23, LOAD_DEFAULT, fd(512, 1030));
if (key_size==8) return wt(192, 15, LOAD_DEFAULT, fd(364, 1085));
if (key_size==16) return wt(128, 23, LOAD_DEFAULT, fd(364, 1050));
if (key_size==1) return wt(128,23,LOAD_DEFAULT,1105);
if (key_size==2) return wt(128,23,LOAD_DEFAULT,1190);
if (key_size==4) return wt(128,23,LOAD_DEFAULT,1030);
if (key_size==8) return wt(192,15,LOAD_DEFAULT,1085);
if (key_size==16) return wt(128,23,LOAD_DEFAULT,1050);
}
// =====================================================================
// SM80 — key 1-8B × value 1-8B (primitive)
// =====================================================================
// SM80
if (pv) {
// key=1B SM80
if (key_size==1 && value_size==1) return dr(128, 12, LOAD_DEFAULT, nd(795));
if (key_size==1 && value_size==2) return wt(288, 12, LOAD_DEFAULT, nd(825));
if (key_size==1 && value_size==4) return wt(256, 15, LOAD_DEFAULT, nd(640));
if (key_size==1 && value_size==8) return wt(192, 10, LOAD_DEFAULT, fd(124, 1040));
// key=2B SM80
if (key_size==2 && value_size==1) return dr(256, 8, LOAD_DEFAULT, nd(1070));
if (key_size==2 && value_size==2) return wt(320, 14, LOAD_DEFAULT, nd(625));
if (key_size==2 && value_size==4) return wt(256, 15, LOAD_DEFAULT, nd(1055));
if (key_size==2 && value_size==8) return wt(160, 17, LOAD_DEFAULT, fd(160, 695));
// key=4B SM80
if (key_size==4 && value_size==1) return dr(128, 12, LOAD_DEFAULT, nd(1130));
if (key_size==4 && value_size==2) return wt(256, 12, LOAD_DEFAULT, nd(1130));
if (key_size==4 && value_size==4) return wt(256, 15, LOAD_DEFAULT, nd(1140));
if (key_size==4 && value_size==8) return wt(256, 9, LOAD_DEFAULT, fd(888, 635));
// key=8B SM80
if (key_size==8 && value_size==1) return wt(128, 11, LOAD_DEFAULT, nd(1120));
if (key_size==8 && value_size==2) return wt(256, 10, LOAD_DEFAULT, nd(1115));
if (key_size==8 && value_size==4) return wt(224, 13, LOAD_DEFAULT, fd(24, 1060));
if (key_size==8 && value_size==8) return wt(224, 10, LOAD_DEFAULT, nd(1160));
if (key_size==1 && value_size==1) return dr(128,12,LOAD_DEFAULT,795);
if (key_size==1 && value_size==2) return wt(288,12,LOAD_DEFAULT,825);
if (key_size==1 && value_size==4) return wt(256,15,LOAD_DEFAULT,640);
if (key_size==1 && value_size==8) return wt(192,10,LOAD_DEFAULT,1040);
if (key_size==2 && value_size==1) return dr(256,8,LOAD_DEFAULT,1070);
if (key_size==2 && value_size==2) return wt(320,14,LOAD_DEFAULT,625);
if (key_size==2 && value_size==4) return wt(256,15,LOAD_DEFAULT,1055);
if (key_size==2 && value_size==8) return wt(160,17,LOAD_DEFAULT,695);
if (key_size==4 && value_size==1) return dr(128,12,LOAD_DEFAULT,1130);
if (key_size==4 && value_size==2) return wt(256,12,LOAD_DEFAULT,1130);
if (key_size==4 && value_size==4) return wt(256,15,LOAD_DEFAULT,1140);
if (key_size==4 && value_size==8) return wt(256,9,LOAD_DEFAULT,635);
if (key_size==8 && value_size==1) return wt(128,11,LOAD_DEFAULT,1120);
if (key_size==8 && value_size==2) return wt(256,10,LOAD_DEFAULT,1115);
if (key_size==8 && value_size==4) return wt(224,13,LOAD_DEFAULT,1060);
if (key_size==8 && value_size==8) return wt(224,10,LOAD_DEFAULT,1160);
}
// SM80 key=16B (accum_is_primitive)
if (key_size==16 && accum_is_primitive) {
if (value_size==1) return wt(192, 7, LOAD_DEFAULT, fd(144, 1120));
if (value_size==2) return wt(192, 7, LOAD_DEFAULT, fd(364, 780));
if (value_size==4) return wt(256, 7, LOAD_DEFAULT, nd(1170));
if (value_size==8) return wt(128, 15, LOAD_DEFAULT, nd(1030));
if (value_size==1) return wt(192,7,LOAD_DEFAULT,1120);
if (value_size==2) return wt(192,7,LOAD_DEFAULT,780);
if (value_size==4) return wt(256,7,LOAD_DEFAULT,1170);
if (value_size==8) return wt(128,15,LOAD_DEFAULT,1030);
}
// SM80 int128 values
if (value_size==16) {
if (key_size==1) return wt(128, 19, LOAD_DEFAULT, nd(1095));
if (key_size==2) return wt(160, 14, LOAD_DEFAULT, nd(1105));
if (key_size==4) return wt(128, 17, LOAD_DEFAULT, nd(1100));
if (key_size==8) return wt(320, 8, LOAD_DEFAULT, nd(220));
if (key_size==16) return wt(128, 15, LOAD_DEFAULT, nd(1160));
if (key_size==1) return wt(128,19,LOAD_DEFAULT,1095);
if (key_size==2) return wt(160,14,LOAD_DEFAULT,1105);
if (key_size==4) return wt(128,17,LOAD_DEFAULT,1100);
if (key_size==8) return wt(320,8,LOAD_DEFAULT,220);
if (key_size==16) return wt(128,15,LOAD_DEFAULT,1160);
}
// =====================================================================
// Default fallback
// =====================================================================
// Default
int mx = key_size > value_size ? key_size : value_size;
int combined = key_size + value_size;
int ipt = (mx <= 8) ? 9 : (9 * 8 / combined);
if (ipt < 1) ipt = 1; if (ipt > 9) ipt = 9;
return wt(256, ipt, LOAD_DEFAULT, fd(350, 450));
int ipt = (mx <= 8) ? 9 : (9*8/(key_size+value_size));
if (ipt<1) ipt=1; if (ipt>9) ipt=9;
return wt(256, ipt, LOAD_DEFAULT, 450);
}
constexpr ScanByKeyPolicy operator()(const hardware_capability& hw) const {