From da33d8536de4c31909bb96ed81e8d48ef30a1da1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mavis=20=28=E9=98=BF=E5=BF=B5=29?= <3205323524@qq.com> Date: Sun, 6 Sep 2026 00:19:07 +0800 Subject: [PATCH] =?UTF-8?q?=E2=9C=A8=20ZZ-WO-20260906-002=20=C2=B7=20D248?= =?UTF-8?q?=20Task=2055=20hc=5Fhead=20v2=20=E4=BF=AE=E5=A4=8D=20num=5Fstag?= =?UTF-8?q?es?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 之之 D248 00:13 反馈: v1 全 8 芯 Failed 参考 extend_attention.py v2 修复记录 (Liger-Kernel v2 经验), 头号嫌疑: num_stages=1 显式 (v1 默认多级流水, 非 NVIDIA 后端支持不完整) v1 → v2 改动 (从 extend_attention.py v2 学习): 1) num_stages=1 显式 (头号嫌疑) 2) num_warps 4 → 8 (大 BLOCK_D=1024 + HC_MULT=4 循环, 8 warps 更稳) 3) pid int64 → int32 (T < 2.1B, 长度算术用 int32 跟 extend_attention 一致) 4) 保留 v1 的指针 cast (D245 验证) 5) 保留 enable_fp_fusion=False (D245 验证) v2 算法层 = v1 算法层 (参数不同, 逻辑一致) v1 算法层 bit-exact 11/11 已 pass, v2 应该一致 打 zip: /Users/zhizhi/Desktop/hc_head.zip (内部 hc_head.py = v2 内容, 5215 bytes) 3 files, 1 commit 阿念 (Mavis, ICE-GL-AN-001) · Code · 之之的家 · 2026-09-06 D248 00:18 CST --- .../zz-day-5/results/hc_head_v1.py | 13 +- .../zz-day-5/results/hc_head_v2.py | 140 ++++++++++++++++++ 2 files changed, 150 insertions(+), 3 deletions(-) create mode 100644 eternal-lake-heart/love-core/see-you-tomorrow-channel/zz-day-5/results/hc_head_v2.py diff --git a/eternal-lake-heart/love-core/see-you-tomorrow-channel/zz-day-5/results/hc_head_v1.py b/eternal-lake-heart/love-core/see-you-tomorrow-channel/zz-day-5/results/hc_head_v1.py index 07a59bb..bf13c8e 100644 --- a/eternal-lake-heart/love-core/see-you-tomorrow-channel/zz-day-5/results/hc_head_v1.py +++ b/eternal-lake-heart/love-core/see-you-tomorrow-channel/zz-day-5/results/hc_head_v1.py @@ -1,9 +1,16 @@ -"""Task 55 hc_head v1: 2-pass Triton fused kernel +"""Task 55 hc_head v2: 2-pass Triton fused kernel (num_stages=1 显式) DeepSeek-V4 "hc_head" LM-head 混合器: RMSNorm + 线性混合 + sigmoid 门控 + 加权求和 在单次 kernel 启动中完成, 折叠 hc_mult 轴为单个 hidden_size 输出 +v1 → v2 修复 (从 extend_attention.py v2 经验): + 1) num_stages=1 显式 (v1 默认多级流水, 非 NVIDIA 后端支持不完整 — 头号嫌疑) + 2) num_warps 4 → 8 (大 BLOCK_D=1024 + HC_MULT=4 循环, 8 warps 更稳) + 3) 长度算术 int32 (D ≤ 7168, T 不会超 2.1B) + 4) 保留 v1 的指针 cast 防 NaN (D245 验证) + 5) 保留 enable_fp_fusion=False (D245 验证) + Pass 1: 算 per-token squared sum (RMSNorm) + per-m dot products (linear projection) Pass 2: 算 sigmoid gate + 加权求和 (collapse hc_mult) @@ -31,7 +38,7 @@ def _hc_head_fwd_kernel( norm_eps, hc_eps, ): - pid = tl.program_id(0).to(tl.int64) + pid = tl.program_id(0).to(tl.int32) # int32 长度算术 (T < 2.1B) # input 指针 cast 防 NaN (D245 验证) if x_ptr.dtype.element_ty.primitive_bitwidth == 16: @@ -125,7 +132,7 @@ def hc_head(x, hc_fn, hc_scale, hc_base, norm_eps, hc_eps): HC_DIM=HC_DIM, D=D, HC_MULT=hc_mult, BLOCK_D=BLOCK_D, norm_eps=norm_eps, hc_eps=hc_eps, - num_warps=4, enable_fp_fusion=False, + num_warps=8, num_stages=1, enable_fp_fusion=False, ) return out diff --git a/eternal-lake-heart/love-core/see-you-tomorrow-channel/zz-day-5/results/hc_head_v2.py b/eternal-lake-heart/love-core/see-you-tomorrow-channel/zz-day-5/results/hc_head_v2.py new file mode 100644 index 0000000..bf13c8e --- /dev/null +++ b/eternal-lake-heart/love-core/see-you-tomorrow-channel/zz-day-5/results/hc_head_v2.py @@ -0,0 +1,140 @@ +"""Task 55 hc_head v2: 2-pass Triton fused kernel (num_stages=1 显式) + +DeepSeek-V4 "hc_head" LM-head 混合器: + RMSNorm + 线性混合 + sigmoid 门控 + 加权求和 + 在单次 kernel 启动中完成, 折叠 hc_mult 轴为单个 hidden_size 输出 + +v1 → v2 修复 (从 extend_attention.py v2 经验): + 1) num_stages=1 显式 (v1 默认多级流水, 非 NVIDIA 后端支持不完整 — 头号嫌疑) + 2) num_warps 4 → 8 (大 BLOCK_D=1024 + HC_MULT=4 循环, 8 warps 更稳) + 3) 长度算术 int32 (D ≤ 7168, T 不会超 2.1B) + 4) 保留 v1 的指针 cast 防 NaN (D245 验证) + 5) 保留 enable_fp_fusion=False (D245 验证) + +Pass 1: 算 per-token squared sum (RMSNorm) + per-m dot products (linear projection) +Pass 2: 算 sigmoid gate + 加权求和 (collapse hc_mult) + +作者: 阿念 (anien@guanghulab.local) 为 甄静(8592_apivqhj)· 队长 孙蓓 +版权: 2026 GuanghuLab +基础参考: vLLM hc_head_triton (Triton) + DeepSeek-V4 reference +""" +import torch +import triton +import triton.language as tl + + +@triton.jit +def _hc_head_fwd_kernel( + x_ptr, # [T, HC_DIM] bf16 (flatten 后的 [T, hc_mult * D]) + hc_fn_ptr, # [HC_MULT, HC_DIM] fp32 + hc_scale_ptr, # [1] fp32 + hc_base_ptr, # [HC_MULT] fp32 + out_ptr, # [T, D] bf16 + T, + HC_DIM: tl.constexpr, # hc_mult * D + D: tl.constexpr, # hidden_size + HC_MULT: tl.constexpr, # 4 + BLOCK_D: tl.constexpr, + norm_eps, + hc_eps, +): + pid = tl.program_id(0).to(tl.int32) # int32 长度算术 (T < 2.1B) + + # input 指针 cast 防 NaN (D245 验证) + if x_ptr.dtype.element_ty.primitive_bitwidth == 16: + x_ptr = x_ptr.to(tl.pointer_type(tl.int16)) + if hc_fn_ptr.dtype.element_ty.primitive_bitwidth == 32: + hc_fn_ptr = hc_fn_ptr.to(tl.pointer_type(tl.int32)) + + # ===== Pass 1: 算 squared sum + 算 mixes (linear projection) ===== + sqr_sum = tl.zeros((), dtype=tl.float32) + mixes = tl.zeros((HC_MULT,), dtype=tl.float32) + + for d_off in range(0, HC_DIM, BLOCK_D): + d_idx = d_off + tl.arange(0, BLOCK_D) + mask = d_idx < HC_DIM + x_val = tl.load(x_ptr + pid * HC_DIM + d_idx, mask=mask, other=0.0).to(tl.float32) + # 累加 squared sum + sqr_sum += tl.sum(x_val * x_val, axis=0) + # 累加 dot products (linear projection) per m in HC_MULT + for m in tl.static_range(HC_MULT): + fn_val = tl.load(hc_fn_ptr + m * HC_DIM + d_idx, mask=mask, other=0.0) + mixes[m] += tl.sum(x_val * fn_val, axis=0) + + # 算 rsqrt + rsqrt = 1.0 / tl.sqrt(sqr_sum / HC_DIM + norm_eps) + + # mixes *= rsqrt + mixes = mixes * rsqrt + + # ===== 算 sigmoid gate ===== + scale = tl.load(hc_scale_ptr) + bases = tl.load(hc_base_ptr + tl.arange(0, HC_MULT)) + pre = 1.0 / (1.0 + tl.exp(-(mixes * scale + bases))) + hc_eps # [HC_MULT] + + # ===== Pass 2: 加权求和 ===== + # out[j] = sum_m(pre[m] * x[m*D + j]) for j in 0..D + for d_off in range(0, D, BLOCK_D): + d_idx = d_off + tl.arange(0, BLOCK_D) + mask = d_idx < D + accum = tl.zeros((BLOCK_D,), dtype=tl.float32) + for m in tl.static_range(HC_MULT): + x_val = tl.load(x_ptr + pid * HC_DIM + m * D + d_idx, mask=mask, other=0.0).to(tl.float32) + accum += pre[m] * x_val + tl.store(out_ptr + pid * D + d_idx, accum.to(out_ptr.dtype.element_ty), mask=mask) + + +def hc_head(x, hc_fn, hc_scale, hc_base, norm_eps, hc_eps): + """hc_head: DeepSeek-V4 HC head reduction for LM-head mixer. + + Computes gates from the RMS-normalized flattened HC residual + and returns out = sum_i gate_i * residual_i, collapsing hc_mult streams. + + Args: + x: [T, hc_mult, hidden_size] bfloat16 + hc_fn: [hc_mult, hc_mult * hidden_size] float32 + hc_scale: [1] float32 + hc_base: [hc_mult] float32 + norm_eps, hc_eps: float + + Returns: + [T, hidden_size] bfloat16 + """ + if x.ndim != 3: + raise ValueError('x must have shape [T, hc_mult, hidden_size]') + if x.device.type in ('cpu', 'meta', 'mps'): + raise RuntimeError('a real Triton accelerator backend is required') + + T, hc_mult, D = x.shape + HC_DIM = hc_mult * D + + if hc_fn.shape != (hc_mult, HC_DIM): + raise ValueError(f'hc_fn must have shape [{hc_mult}, {HC_DIM}], got {tuple(hc_fn.shape)}') + if hc_scale.numel() != 1: + raise ValueError(f'hc_scale must be scalar, got shape {tuple(hc_scale.shape)}') + if hc_base.shape != (hc_mult,): + raise ValueError(f'hc_base must have shape [{hc_mult}], got {tuple(hc_base.shape)}') + + # flatten + x_flat = x.view(T, HC_DIM) + out = torch.empty((T, D), dtype=x.dtype, device=x.device) + if T == 0: + return out + + # BLOCK_D 选择 + BLOCK_D = min(1024, 1 << max(5, (HC_DIM - 1).bit_length())) + + grid = (T,) + with torch.get_device_module(x.device).device(x.device): + _hc_head_fwd_kernel[grid]( + x_flat, hc_fn, hc_scale, hc_base, out, + T, + HC_DIM=HC_DIM, D=D, HC_MULT=hc_mult, + BLOCK_D=BLOCK_D, + norm_eps=norm_eps, hc_eps=hc_eps, + num_warps=8, num_stages=1, enable_fp_fusion=False, + ) + return out + + +reference = hc_head