From b50bd2dfd5dad12cb3da3e5a8339fa0689341b81 Mon Sep 17 00:00:00 2001 From: muh-bot Date: Wed, 5 Aug 2026 03:20:39 +0000 Subject: [PATCH] [gen_patch] fix critical struct selection: dispatch by kernel data type MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit CCCL policy_selector dispatches by (accum_size, type_t, offset_size). gen_patch was picking first non-default struct → bi100_plus_accum1_o4 (int8 items=32) for reduce. paged_attention uses float32 scores → correct struct is bi100_plus_float32_o4 (items=24). Before: 512*32*4=65536 > 49152 SMEM → crash After: 512*24*4=49152 = 100% SMEM → correct SCAN_BLOCK_SIZE: 512→384 (bench_bi100.py ipt=22 tpb=384 dcid=0) --- muh/gen_patch.py | 70 +++++++++++++++++++++++++++++++++--------------- 1 file changed, 48 insertions(+), 22 deletions(-) diff --git a/muh/gen_patch.py b/muh/gen_patch.py index 082f9957..98c91343 100644 --- a/muh/gen_patch.py +++ b/muh/gen_patch.py @@ -268,29 +268,55 @@ def generate_patches(header_dir): summary.append(f"SKIP {algo}: no bi100_* structs and no inline values found") continue - # Select the most relevant struct for vllm's primary data path. - # vllm's paged_attention score accumulator is always float32 (4 bytes), - # so we prefer bi100_*float32* or bi100_*accum4* structs. - # Fallback priority: float32 > accum2 (fp16 KV) > first non-default > first. - primary = None - preference_order = ['float32', 'accum4', 'accum2', 'int32'] - for pref in preference_order: - for name, fields in structs: - if pref in name and 'det' not in name and 'default' not in name: - primary = (name, fields) - break - if primary: - break - if primary is None: - for name, fields in structs: - if 'default' not in name and 'det' not in name: - primary = (name, fields) - break - if primary is None: - primary = structs[0] + # Select the struct that matches each vllm kernel's data type. + # + # CCCL's policy_selector dispatches by (accum_size, type_t, offset_size). + # gen_patch must do the same: when injecting into paged_attention + # (float32 scores), use bi100_plus_float32_o4, not bi100_plus_accum1_o4. + # + # The VLLM_KERNEL_MAP in muh_kernel_map.py defines each kernel's + # data_types. This mapping encodes the primary data type per algorithm: + ALGO_PRIMARY_TYPE = { + 'reduce': ('float32', 4), # paged_attention scores + 'scan': ('float32', 4), # softmax denominator + 'topk': ('float32', 4), # logits + 'transform': ('float16', 2), # activations (SiLU, RMSNorm input) + 'batch_memcpy': ('float16', 2), # KV cache blocks + 'for': ('float16', 2), # RoPE + } - pname, pfields = primary - summary.append(f"READ {algo}: {pname} → {pfields}") + target_type, target_size = ALGO_PRIMARY_TYPE.get(algo, ('float32', 4)) + + # Score each struct by match quality + def struct_score(name, fields): + score = 0 + name_lower = name.lower() + # Exact type name match (best) + if target_type.replace('float', 'f') in name_lower or target_type in name_lower: + score += 100 + # Accum/type size match in name (e.g. "_4B_", "_accum4_", "float32") + size_tags = [f'_{target_size}B', f'_accum{target_size}', f'float{target_size*8}'] + for tag in size_tags: + if tag.lower() in name_lower: + score += 50 + # Offset size 4 preferred (most common in vllm) + if '_o4' in name_lower: + score += 10 + # Penalize 'default' and 'det' (deterministic) structs + if 'default' in name_lower: + score -= 200 + if 'det' in name_lower: + score -= 50 + # Penalize 1-byte type structs for float32 targets + if target_size >= 4 and ('_1B' in name or 'accum1' in name_lower): + score -= 100 + return score + + scored = [(struct_score(n, f), n, f) for n, f in structs] + scored.sort(key=lambda x: -x[0]) + _, pname, pfields = scored[0] + + summary.append(f"READ {algo}: {pname} → {pfields} (target: {target_type})") for field_name, value in pfields.items(): key = (algo, field_name)