dsv4.1: mHC computation and compensated projections (#39664)

Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
  Co-authored-by: BBuf <1182563586@qq.com>
  Co-authored-by: zijiexia <37504505+zijiexia@users.noreply.github.com>
This commit is contained in:
Liangsheng Yin
2026-09-16 17:11:55 -07:00
committed by GitHub
parent 46ae84df15
commit 408d2334c3
+521 -1
View File
@@ -18,6 +18,7 @@ from sglang.srt.environ import envs
from sglang.srt.layers.attention.dsa.utils import is_dsa_prefill_cp_interleave
from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.utils.common import strict_contiguous
from sglang.srt.runtime_context import get_platform
from sglang.srt.utils.common import is_gfx1250_supported
logger = logging.getLogger(__name__)
@@ -362,13 +363,21 @@ def hc_split_sinkhorn(
sinkhorn_iters: int = 20,
eps: float = 1e-6,
):
b, s, _ = mixes.size()
if b * s == 0:
# DP attention's idle forward carries no tokens, and every backend below
# derives a CUDA grid from the token count; a zero-sized grid is rejected.
return (
mixes.new_empty(b, s, hc_mult),
mixes.new_empty(b, s, hc_mult),
mixes.new_empty(b, s, hc_mult, hc_mult),
)
if is_gfx1250_supported():
# TileLang's CK-backed addressing doesn't compile on gfx1250; use the
# Triton port. _hc_split_sinkhorn_torch is kept as a reference fallback.
return _hc_split_sinkhorn_triton(
mixes, hc_scale, hc_base, hc_mult, sinkhorn_iters, eps
)
b, s, _ = mixes.size()
pre = mixes.new_empty(b, s, hc_mult)
post = mixes.new_empty(b, s, hc_mult)
comb = mixes.new_empty(b, s, hc_mult, hc_mult)
@@ -2027,6 +2036,517 @@ def _hc_combine_kernel(
tl.store(y_ptr + pid_m * y_stride_m + offs_h, acc, mask=mask)
@triton.jit
def _hc_mix_stats_partial_kernel(
x_ptr,
w_ptr,
part_mix_ptr,
part_sq_ptr,
M,
K,
x_stride_m,
w_stride_n,
MIX: tl.constexpr,
MIX_PAD: tl.constexpr,
NUM_SLICES: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_K: tl.constexpr,
DOT_PRECISION: tl.constexpr,
):
"""Mixing dot products and row sum of squares over one K slice."""
pid_m = tl.program_id(0)
pid_s = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = tl.arange(0, MIX_PAD)
mask_m = offs_m < M
mask_n = offs_n < MIX
k_per_slice = K // NUM_SLICES
k_start = pid_s * k_per_slice
acc = tl.zeros([BLOCK_M, MIX_PAD], dtype=tl.float32)
sq = tl.zeros([BLOCK_M], dtype=tl.float32)
for kb in range(0, k_per_slice, BLOCK_K):
offs_k = k_start + kb + tl.arange(0, BLOCK_K)
mask_k = offs_k < k_start + k_per_slice
x_tile = tl.load(
x_ptr + offs_m[:, None] * x_stride_m + offs_k[None, :],
mask=mask_m[:, None] & mask_k[None, :],
other=0.0,
).to(tl.float32)
w_tile = tl.load(
w_ptr + offs_n[None, :] * w_stride_n + offs_k[:, None],
mask=mask_n[None, :] & mask_k[:, None],
other=0.0,
).to(tl.float32)
acc += tl.dot(x_tile, w_tile, input_precision=DOT_PRECISION)
sq += tl.sum(x_tile * x_tile, axis=1)
tl.store(
part_mix_ptr + (pid_s * M + offs_m[:, None]) * MIX + offs_n[None, :],
acc,
mask=mask_m[:, None] & mask_n[None, :],
)
tl.store(part_sq_ptr + pid_s * M + offs_m, sq, mask=mask_m)
@triton.jit
def _hc_mix_stats_reduce_kernel(
part_mix_ptr,
part_sq_ptr,
mixes_ptr,
M,
inv_k,
eps,
MIX: tl.constexpr,
MIX_PAD: tl.constexpr,
NUM_SLICES: tl.constexpr,
BLOCK_M: tl.constexpr,
):
"""Sum the NUM_SLICES partials in slice order and apply the rms scaling."""
pid_m = tl.program_id(0)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = tl.arange(0, MIX_PAD)
mask_m = offs_m < M
mask_n = offs_n < MIX
acc = tl.zeros([BLOCK_M, MIX_PAD], dtype=tl.float32)
sq = tl.zeros([BLOCK_M], dtype=tl.float32)
for s in tl.static_range(NUM_SLICES):
acc += tl.load(
part_mix_ptr + (s * M + offs_m[:, None]) * MIX + offs_n[None, :],
mask=mask_m[:, None] & mask_n[None, :],
other=0.0,
)
sq += tl.load(part_sq_ptr + s * M + offs_m, mask=mask_m, other=0.0)
rsqrt = 1.0 / tl.sqrt(sq * inv_k + eps)
tl.store(
mixes_ptr + offs_m[:, None] * MIX + offs_n[None, :],
acc * rsqrt[:, None],
mask=mask_m[:, None] & mask_n[None, :],
)
# K slicing, BLOCK_K and dot precision must stay independent of M;
# a row must produce the same bits alone and in any batch routed to this backend.
_HC_MIX_SLICE_CHOICES = (80, 64, 40, 32, 16, 8, 4, 2, 1)
_HC_MIX_BLOCK_M = 32
_HC_MIX_BLOCK_K = 64
_HC_MIX_NUM_WARPS = 4
_HC_MIX_DOT_PRECISION = "tf32x3"
# num_stages only reorders memory issue, not arithmetic; 2 is enough to cover the
# short k_per_slice loop (K=20480 gives 80 slices, i.e. 4 BLOCK_K tiles per CTA).
_HC_MIX_NUM_STAGES = 2
# BLOCK_M 8/16/32 preserve each row's K reduction order; thresholds were measured on GB300.
# Keep BLOCK_M below 64, where Triton lowers tf32x3 to plain TF32 and changes rounding.
_HC_MIX_BLOCK_M_SMALL = 8
_HC_MIX_BLOCK_M_MID = 16
_HC_MIX_MID_MAX_M = 2048
def _block_m_for(m: int) -> int:
"""Row-tile choices preserve each row's arithmetic and may depend on M."""
if m <= _HC_MIX_BLOCK_M_SMALL:
return _HC_MIX_BLOCK_M_SMALL
if m <= _HC_MIX_MID_MAX_M:
return _HC_MIX_BLOCK_M_MID
return _HC_MIX_BLOCK_M
def _num_stages_for(m: int, k: int) -> int:
# GB300 verify batches benefit from a smaller shared-memory footprint.
if get_platform().is_blackwell and k == 20480 and 64 <= m <= 384:
return 1
return _HC_MIX_NUM_STAGES
def _num_slices_for(k: int) -> int:
"""Slice count depends only on K, never on batch size M."""
blocks = k // _HC_MIX_BLOCK_K
assert k % _HC_MIX_BLOCK_K == 0, k
for n in _HC_MIX_SLICE_CHOICES:
if blocks % n == 0:
return n
return 1
def hc_mix_stats(x_flat: torch.Tensor, hc_fn: torch.Tensor, eps: float) -> torch.Tensor:
"""Batch-invariant F.linear(x_flat.float(), hc_fn) * rsqrt(mean(x_flat^2) + eps).
x_flat is [M, K] in any float dtype; hc_fn is [MIX, K] fp32; returns [M, MIX] fp32.
K slicing and reduction order are independent of M, so each row is bitwise
identical whether computed alone or in any batch routed to this backend.
"""
assert x_flat.dim() == 2 and hc_fn.dim() == 2
assert x_flat.stride(1) == 1 and hc_fn.stride(1) == 1
assert hc_fn.dtype == torch.float32
m, k = x_flat.shape
mix = hc_fn.shape[0]
assert hc_fn.shape[1] == k
num_slices = _num_slices_for(k)
mix_pad = max(16, triton.next_power_of_2(mix))
part_mix = torch.empty(
(num_slices, m, mix), dtype=torch.float32, device=x_flat.device
)
part_sq = torch.empty((num_slices, m), dtype=torch.float32, device=x_flat.device)
mixes = torch.empty((m, mix), dtype=torch.float32, device=x_flat.device)
if m == 0:
return mixes
block_m = _block_m_for(m)
grid_m = triton.cdiv(m, block_m)
_hc_mix_stats_partial_kernel[(grid_m, num_slices)](
x_flat,
hc_fn,
part_mix,
part_sq,
m,
k,
x_flat.stride(0),
hc_fn.stride(0),
MIX=mix,
MIX_PAD=mix_pad,
NUM_SLICES=num_slices,
BLOCK_M=block_m,
BLOCK_K=_HC_MIX_BLOCK_K,
DOT_PRECISION=_HC_MIX_DOT_PRECISION,
num_warps=_HC_MIX_NUM_WARPS,
num_stages=_num_stages_for(m, k),
)
_hc_mix_stats_reduce_kernel[(grid_m,)](
part_mix,
part_sq,
mixes,
m,
1.0 / k,
eps,
MIX=mix,
MIX_PAD=mix_pad,
NUM_SLICES=num_slices,
BLOCK_M=block_m,
num_warps=4,
)
return mixes
@triton.jit
def _hc_mix_reduce_sinkhorn_kernel(
part_mix_ptr,
part_sq_ptr,
scale_ptr,
base_ptr,
pre_ptr,
post_ptr,
comb_ptr,
m,
inv_k,
rms_eps,
MIX: tl.constexpr,
HC: tl.constexpr,
NUM_SLICES: tl.constexpr,
ITERS: tl.constexpr,
EPS: tl.constexpr,
part_mix_residual_ptr=None,
):
"""One CTA per row keeps the sinkhorn reductions two-dimensional."""
row = tl.program_id(0)
if row >= m:
return
j = tl.arange(0, HC)
jj = j[:, None]
kk = j[None, :]
a_pre = tl.zeros([HC], dtype=tl.float32)
a_post = tl.zeros([HC], dtype=tl.float32)
a_comb = tl.zeros([HC, HC], dtype=tl.float32)
sq = tl.zeros([], dtype=tl.float32)
for s in tl.static_range(NUM_SLICES):
off = (s * m + row) * MIX
v_pre = tl.load(part_mix_ptr + off + j)
v_post = tl.load(part_mix_ptr + off + HC + j)
v_comb = tl.load(part_mix_ptr + off + 2 * HC + jj * HC + kk)
if part_mix_residual_ptr is not None:
v_pre += tl.load(part_mix_residual_ptr + off + j)
v_post += tl.load(part_mix_residual_ptr + off + HC + j)
v_comb += tl.load(part_mix_residual_ptr + off + 2 * HC + jj * HC + kk)
a_pre += v_pre
a_post += v_post
a_comb += v_comb
sq += tl.load(part_sq_ptr + s * m + row)
rsqrt = 1.0 / tl.sqrt(sq * inv_k + rms_eps)
s0 = tl.load(scale_ptr + 0)
s1 = tl.load(scale_ptr + 1)
s2 = tl.load(scale_ptr + 2)
pre = tl.sigmoid(a_pre * rsqrt * s0 + tl.load(base_ptr + j)) + EPS
tl.store(pre_ptr + row * HC + j, pre)
post = 2.0 * tl.sigmoid(a_post * rsqrt * s1 + tl.load(base_ptr + HC + j))
tl.store(post_ptr + row * HC + j, post)
comb = a_comb * rsqrt * s2 + tl.load(base_ptr + 2 * HC + jj * HC + kk)
comb = tl.exp(comb - tl.max(comb, axis=1)[:, None])
comb = comb / tl.sum(comb, axis=1)[:, None] + EPS
comb = comb / (tl.sum(comb, axis=0)[None, :] + EPS)
for _ in tl.static_range(ITERS - 1):
comb = comb / (tl.sum(comb, axis=1)[:, None] + EPS)
comb = comb / (tl.sum(comb, axis=0)[None, :] + EPS)
tl.store(comb_ptr + row * HC * HC + jj * HC + kk, comb)
def hc_mix_stats_sinkhorn(
x_flat: torch.Tensor,
hc_fn: torch.Tensor,
hc_scale: torch.Tensor,
hc_base: torch.Tensor,
hc_mult: int,
sinkhorn_iters: int,
rms_eps: float,
hc_eps: float,
):
"""Fuse the reduce and sinkhorn stages of hc_mix_stats followed by hc_split_sinkhorn.
The split-K kernel fixes the reduction order; sinkhorn uses the Triton port's
transcendental lowering, which differs from TileLang.
"""
assert x_flat.dim() == 2 and hc_fn.dim() == 2
assert x_flat.stride(1) == 1 and hc_fn.stride(1) == 1
assert hc_fn.dtype == torch.float32
m, k = x_flat.shape
mix = hc_fn.shape[0]
assert mix == (2 + hc_mult) * hc_mult and hc_fn.shape[1] == k
dev = x_flat.device
pre = torch.empty(m, hc_mult, dtype=torch.float32, device=dev)
post = torch.empty(m, hc_mult, dtype=torch.float32, device=dev)
comb = torch.empty(m, hc_mult, hc_mult, dtype=torch.float32, device=dev)
if m == 0:
return pre, post, comb
num_slices = _num_slices_for(k)
mix_pad = max(16, triton.next_power_of_2(mix))
part_mix = torch.empty((num_slices, m, mix), dtype=torch.float32, device=dev)
part_sq = torch.empty((num_slices, m), dtype=torch.float32, device=dev)
block_m = _block_m_for(m)
_hc_mix_stats_partial_kernel[(triton.cdiv(m, block_m), num_slices)](
x_flat,
hc_fn,
part_mix,
part_sq,
m,
k,
x_flat.stride(0),
hc_fn.stride(0),
MIX=mix,
MIX_PAD=mix_pad,
NUM_SLICES=num_slices,
BLOCK_M=block_m,
BLOCK_K=_HC_MIX_BLOCK_K,
DOT_PRECISION=_HC_MIX_DOT_PRECISION,
num_warps=_HC_MIX_NUM_WARPS,
num_stages=_num_stages_for(m, k),
)
_hc_mix_reduce_sinkhorn_kernel[(m,)](
part_mix,
part_sq,
hc_scale.float().contiguous(),
hc_base.float().contiguous(),
pre,
post,
comb,
m,
1.0 / k,
rms_eps,
MIX=mix,
HC=hc_mult,
NUM_SLICES=num_slices,
ITERS=sinkhorn_iters,
EPS=hc_eps,
num_warps=1,
)
return pre, post, comb
# Compensated projections: the fp32 mixing weight is split into components the
# tensor cores take exactly (three bf16 parts, or a tf32 high part and its fp32
# residual), accumulated over a fixed slice count to bound the fp32 error.
_HC_MIX_COMPENSATED_SLICES = 16
_HC_MIX_BF16X3_BLOCK_M = 128
def split_bf16_hc_weight(weight: torch.Tensor):
assert weight.dtype == torch.float32 and weight.is_contiguous()
high = weight.bfloat16()
residual = weight - high.float()
middle = residual.bfloat16()
low = (residual - middle.float()).bfloat16()
return high, middle, low
def split_tf32_hc_weight(weight: torch.Tensor):
assert weight.dtype == torch.float32 and weight.is_contiguous()
high = (weight.view(torch.int32) & -8192).view(torch.float32)
return high, weight - high
@triton.jit
def _hc_mix_stats_bf16x3_kernel(
X,
W_HI,
W_MID,
W_LO,
MIX,
SQ,
M,
K: tl.constexpr,
K_PER_SLICE: tl.constexpr,
MIX_COLS: tl.constexpr,
MIX_PAD: tl.constexpr,
BLOCK_K: tl.constexpr,
BLOCK_M: tl.constexpr,
):
# M stays runtime-valued so variable prefill lengths reuse the same binary.
rows = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
cols = tl.arange(0, MIX_PAD)
start = tl.program_id(1) * K_PER_SLICE
ks = start + tl.arange(0, BLOCK_K)
hi = tl.zeros((BLOCK_M, MIX_PAD), tl.float32)
mid = tl.zeros((BLOCK_M, MIX_PAD), tl.float32)
lo = tl.zeros((BLOCK_M, MIX_PAD), tl.float32)
sq = tl.zeros((BLOCK_M,), tl.float32)
for block in range(K_PER_SLICE // BLOCK_K):
k = ks + block * BLOCK_K
x = tl.load(
X + rows[:, None].to(tl.int64) * K + k[None, :],
rows[:, None] < M,
0,
)
offsets = cols[None, :] * K + k[:, None]
w_hi = tl.load(W_HI + offsets, cols[None, :] < MIX_COLS, 0)
w_mid = tl.load(W_MID + offsets, cols[None, :] < MIX_COLS, 0)
w_lo = tl.load(W_LO + offsets, cols[None, :] < MIX_COLS, 0)
hi = tl.dot(x, w_hi, hi)
mid = tl.dot(x, w_mid, mid)
lo = tl.dot(x, w_lo, lo)
xf = x.to(tl.float32)
sq += tl.sum(xf * xf, 1)
offsets = (tl.program_id(1) * M + rows[:, None]) * MIX_COLS + cols[None, :]
tl.store(
MIX + offsets, (hi + mid) + lo, (rows[:, None] < M) & (cols[None, :] < MIX_COLS)
)
tl.store(SQ + tl.program_id(1) * M + rows, sq, rows < M)
def hc_mix_stats_sinkhorn_bf16x3(
x: torch.Tensor,
weight_parts,
scale: torch.Tensor,
base: torch.Tensor,
sinkhorn_iters: int,
rms_eps: float,
hc_eps: float,
hc_mult: int = 4,
):
m, k = x.shape
mix = (2 + hc_mult) * hc_mult
slices = _HC_MIX_COMPENSATED_SLICES
assert x.is_contiguous() and x.dtype == torch.bfloat16 and 4096 <= m <= 65536
assert k % (slices * _HC_MIX_BLOCK_K) == 0
assert len(weight_parts) == 3
assert all(
w.shape == (mix, k) and w.dtype == torch.bfloat16 and w.is_contiguous()
for w in weight_parts
)
part_mix = torch.empty((slices, m, mix), device=x.device, dtype=torch.float32)
sq = torch.empty((slices, m), device=x.device, dtype=torch.float32)
pre = torch.empty((m, hc_mult), device=x.device, dtype=torch.float32)
post = torch.empty_like(pre)
comb = torch.empty((m, hc_mult, hc_mult), device=x.device, dtype=torch.float32)
_hc_mix_stats_bf16x3_kernel[(triton.cdiv(m, _HC_MIX_BF16X3_BLOCK_M), slices)](
x,
*weight_parts,
part_mix,
sq,
m,
K=k,
K_PER_SLICE=k // slices,
MIX_COLS=mix,
MIX_PAD=triton.next_power_of_2(mix),
BLOCK_K=_HC_MIX_BLOCK_K,
BLOCK_M=_HC_MIX_BF16X3_BLOCK_M,
num_warps=4,
num_stages=3,
)
_hc_mix_reduce_sinkhorn_kernel[(m,)](
part_mix,
sq,
scale,
base,
pre,
post,
comb,
m,
1.0 / k,
rms_eps,
MIX=mix,
HC=hc_mult,
NUM_SLICES=slices,
ITERS=sinkhorn_iters,
EPS=hc_eps,
num_warps=1,
)
return pre, post, comb
def hc_mix_stats_sinkhorn_deepgemm(
x_flat: torch.Tensor,
weight_parts,
hc_scale: torch.Tensor,
hc_base: torch.Tensor,
sinkhorn_iters: int,
rms_eps: float,
hc_eps: float,
hc_mult: int = 4,
):
from sglang.srt.layers.deep_gemm_wrapper.entrypoint import tf32_hc_prenorm_gemm
assert x_flat.dtype == torch.bfloat16 and x_flat.is_contiguous()
m, k = x_flat.shape
mix = (2 + hc_mult) * hc_mult
slices = _HC_MIX_COMPENSATED_SLICES
high, low = weight_parts
assert high.shape == low.shape == (mix, k)
dev = x_flat.device
pre = torch.empty((m, hc_mult), dtype=torch.float32, device=dev)
post = torch.empty_like(pre)
comb = torch.empty((m, hc_mult, hc_mult), dtype=torch.float32, device=dev)
if m == 0:
return pre, post, comb
mix_hi = torch.empty((slices, m, mix), dtype=torch.float32, device=dev)
mix_lo = torch.empty_like(mix_hi)
sq = torch.empty((slices, m), dtype=torch.float32, device=dev)
tf32_hc_prenorm_gemm(x_flat, high, mix_hi, sq, slices)
# sq depends only on x_flat, so the second projection recomputes the same
# values; overwriting it avoids a throwaway (slices, m) scratch, 4 MiB at m=65536.
tf32_hc_prenorm_gemm(x_flat, low, mix_lo, sq, slices)
_hc_mix_reduce_sinkhorn_kernel[(m,)](
mix_hi,
sq,
hc_scale,
hc_base,
pre,
post,
comb,
m,
1.0 / k,
rms_eps,
MIX=mix,
HC=hc_mult,
NUM_SLICES=slices,
ITERS=sinkhorn_iters,
EPS=hc_eps,
part_mix_residual_ptr=mix_lo,
num_warps=1,
)
return pre, post, comb
def hc_combine(
x_flat: torch.Tensor, pre: torch.Tensor, hc: int, out_dtype: torch.dtype
) -> torch.Tensor: