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:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user