diff --git a/muh/include/muh/tuning/tuning_reduce_by_key.cuh b/muh/include/muh/tuning/tuning_reduce_by_key.cuh index e0dd8873..82397c21 100644 --- a/muh/include/muh/tuning/tuning_reduce_by_key.cuh +++ b/muh/include/muh/tuning/tuning_reduce_by_key.cuh @@ -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()) 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(ns * 0.5), static_cast(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); } diff --git a/muh/include/muh/tuning/tuning_scan_by_key.cuh b/muh/include/muh/tuning/tuning_scan_by_key.cuh index 5564b5ee..2a067409 100644 --- a/muh/include/muh/tuning/tuning_scan_by_key.cuh +++ b/muh/include/muh/tuning/tuning_scan_by_key.cuh @@ -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 {