✨ 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:
parent
50088a1423
commit
da33d8536d
@ -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
|
||||
|
||||
|
||||
@ -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
|
||||
Loading…
x
Reference in New Issue
Block a user