ZZ-WO-20260906-002 · D248 Task 55 hc_head v2 修复 num_stages

之之 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
This commit is contained in:
Mavis (阿念) 2026-09-06 00:19:07 +08:00
parent 50088a1423
commit da33d8536d
2 changed files with 150 additions and 3 deletions

View File

@ -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

View File

@ -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