[DeepSeek-V4] Add mhc_fused_post_pre kernel (#25976)
Co-authored-by: Qichao Li <liqichao@baidu.com>
This commit is contained in:
@@ -657,6 +657,7 @@ class Envs:
|
||||
SGLANG_OPT_USE_TILELANG_MHC_PRE = EnvBool(True)
|
||||
SGLANG_OPT_USE_TILELANG_MHC_POST = EnvBool(True)
|
||||
SGLANG_OPT_USE_TRITON_FUSED_MHC = EnvBool(True)
|
||||
SGLANG_OPT_FUSE_MHC_POST_PRE = EnvBool(False)
|
||||
SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False)
|
||||
SGLANG_OPT_USE_AITER_INDEXER = EnvBool(False)
|
||||
SGLANG_OPT_USE_JIT_INDEXER_METADATA = EnvBool(True)
|
||||
|
||||
@@ -896,3 +896,500 @@ def mhc_post(
|
||||
residual.shape[-1],
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
@tilelang.jit(
|
||||
pass_configs={
|
||||
tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True,
|
||||
tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True,
|
||||
tilelang.PassConfigKey.TL_PTXAS_REGISTER_USAGE_LEVEL: 10,
|
||||
},
|
||||
)
|
||||
def mhc_fused_post_pre_fma_tilelang(
|
||||
prev_comb_mix,
|
||||
prev_residual,
|
||||
prev_post_mix,
|
||||
hidden_in,
|
||||
pre_fn,
|
||||
mixes_partial_out,
|
||||
sqrsum_partial_out,
|
||||
cur_residual_out,
|
||||
hc: int,
|
||||
hidden_size: int,
|
||||
num_mix_outputs: int,
|
||||
n_thr: int = 256,
|
||||
tile_mix_outputs: int = 1,
|
||||
split_k: int = 1,
|
||||
) -> tilelang.JITKernel:
|
||||
num_tokens = T.dynamic("num_tokens")
|
||||
split_k = T.dynamic("split_k")
|
||||
|
||||
hidden_per_split = (hidden_size + split_k - 1) // split_k
|
||||
num_mix_output_tiles = (num_mix_outputs + tile_mix_outputs - 1) // tile_mix_outputs
|
||||
|
||||
prev_comb_mix: T.Tensor((num_tokens, hc, hc), T.float32)
|
||||
prev_residual: T.Tensor((num_tokens, hc, hidden_size), T.bfloat16)
|
||||
prev_post_mix: T.Tensor((num_tokens, hc), T.float32)
|
||||
hidden_in: T.Tensor((num_tokens, hidden_size), T.bfloat16)
|
||||
pre_fn: T.Tensor((num_mix_outputs, hc, hidden_size), T.float32)
|
||||
|
||||
mixes_partial_out: T.Tensor((split_k, num_tokens, num_mix_outputs), T.float32)
|
||||
sqrsum_partial_out: T.Tensor((split_k, num_tokens), T.float32)
|
||||
cur_residual_out: T.Tensor((num_tokens, hc, hidden_size), T.bfloat16)
|
||||
|
||||
hidden_iters_per_thread = (hidden_per_split + n_thr - 1) // n_thr
|
||||
num_warps = n_thr // 32
|
||||
|
||||
ENABLE_PDL = is_arch_support_pdl()
|
||||
|
||||
# CTA assignment:
|
||||
# token_idx : this CTA handles one token.
|
||||
# mix_output_tile_idx : this CTA handles a small tile of mix output columns.
|
||||
# For HC=4, num_mix_outputs = 24:
|
||||
# [0:4] -> pre logits
|
||||
# [4:8] -> post logits
|
||||
# [8:24] -> comb logits
|
||||
# hidden_split_idx : this CTA handles one split of the hidden dimension.
|
||||
#
|
||||
# Thread assignment inside one CTA:
|
||||
# Each thread owns several hidden positions in this hidden split:
|
||||
# hidden_idx = hidden_split_start + hidden_iter * n_thr + thread_idx
|
||||
#
|
||||
# For each owned hidden_idx, the thread computes:
|
||||
# 1. post result: cur_residual[token, :, hidden_idx]
|
||||
# 2. sqrsum partial for pre RMS
|
||||
# 3. GEMM partial for several mix output columns
|
||||
with T.Kernel(
|
||||
num_tokens,
|
||||
num_mix_output_tiles,
|
||||
split_k,
|
||||
threads=n_thr,
|
||||
) as (token_idx, mix_output_tile_idx, hidden_split_idx):
|
||||
thread_idx = T.get_thread_binding()
|
||||
warp_idx = T.get_warp_idx()
|
||||
lane_idx = T.get_lane_idx()
|
||||
|
||||
warp_partials = T.alloc_shared((num_warps, tile_mix_outputs + 1), T.float32)
|
||||
post_mix_smem = T.alloc_shared((hc,), T.float32)
|
||||
comb_mix_smem = T.alloc_shared((hc, hc), T.float32)
|
||||
|
||||
post_mix_for_token = T.alloc_local((hc,), T.float32)
|
||||
comb_mix_for_token = T.alloc_local((hc, hc), T.float32)
|
||||
|
||||
mix_acc = T.alloc_local((tile_mix_outputs,), T.float32)
|
||||
sqrsum_acc = T.alloc_local((1,), T.float32)
|
||||
cur_residual_values = T.alloc_local((hc,), T.float32)
|
||||
|
||||
T.clear(mix_acc)
|
||||
T.clear(sqrsum_acc)
|
||||
|
||||
hidden_split_start = hidden_split_idx * hidden_per_split
|
||||
|
||||
if ENABLE_PDL:
|
||||
T.pdl_sync()
|
||||
|
||||
# Load post/comb coefficients for this token.
|
||||
#
|
||||
# PyTorch equivalent:
|
||||
# post = prev_post_mix[token_idx] # [HC]
|
||||
# comb = prev_comb_mix[token_idx] # [HC, HC]
|
||||
T.copy(prev_post_mix[token_idx, 0], post_mix_smem)
|
||||
T.copy(prev_comb_mix[token_idx, 0, 0], comb_mix_smem)
|
||||
|
||||
for route_idx in T.unroll(hc):
|
||||
post_mix_for_token[route_idx] = post_mix_smem[route_idx]
|
||||
|
||||
for old_route_idx in T.unroll(hc):
|
||||
for new_route_idx in T.unroll(hc):
|
||||
comb_mix_for_token[old_route_idx, new_route_idx] = comb_mix_smem[
|
||||
old_route_idx, new_route_idx
|
||||
]
|
||||
|
||||
for hidden_iter in T.serial(hidden_iters_per_thread):
|
||||
hidden_idx = hidden_split_start + hidden_iter * n_thr + thread_idx
|
||||
|
||||
if hidden_idx < hidden_size:
|
||||
# Step A: fused post.
|
||||
#
|
||||
# PyTorch equivalent:
|
||||
# cur_residual =
|
||||
# post.unsqueeze(-1) * hidden_in.unsqueeze(1)
|
||||
# + (
|
||||
# comb.unsqueeze(-1)
|
||||
# * prev_residual.unsqueeze(2)
|
||||
# ).sum(dim=1)
|
||||
#
|
||||
# Scalar form for this token and this hidden position:
|
||||
# cur_residual[j, h]
|
||||
# = post[j] * hidden_in[h]
|
||||
# + sum_k comb[k, j] * prev_residual[k, h]
|
||||
for new_route_idx in T.unroll(hc):
|
||||
cur_residual_values[new_route_idx] = (
|
||||
post_mix_for_token[new_route_idx]
|
||||
* hidden_in[token_idx, hidden_idx]
|
||||
)
|
||||
|
||||
for old_route_idx in T.unroll(hc):
|
||||
cur_residual_values[new_route_idx] += (
|
||||
comb_mix_for_token[old_route_idx, new_route_idx]
|
||||
* prev_residual[token_idx, old_route_idx, hidden_idx]
|
||||
)
|
||||
|
||||
# Match the unfused path:
|
||||
# mhc_post writes bf16 residual,
|
||||
# then mhc_pre reads bf16 residual.
|
||||
for route_idx in T.unroll(hc):
|
||||
cur_residual_values[route_idx] = T.bfloat16(
|
||||
cur_residual_values[route_idx]
|
||||
)
|
||||
|
||||
# Step B1: pre sqrsum partial.
|
||||
#
|
||||
# PyTorch equivalent:
|
||||
# x_flat = cur_residual.reshape(T, HC * H).float()
|
||||
# sqrsum = (x_flat * x_flat).sum(dim=-1)
|
||||
#
|
||||
# Only mix_output_tile_idx == 0 writes cur_residual and sqrsum,
|
||||
# otherwise different output-column CTAs would duplicate this work.
|
||||
if mix_output_tile_idx == 0:
|
||||
for route_idx in T.unroll(hc):
|
||||
cur_residual_out[token_idx, route_idx, hidden_idx] = (
|
||||
cur_residual_values[route_idx]
|
||||
)
|
||||
sqrsum_acc[0] += (
|
||||
cur_residual_values[route_idx]
|
||||
* cur_residual_values[route_idx]
|
||||
)
|
||||
|
||||
# Step B2: pre GEMM partial.
|
||||
#
|
||||
# PyTorch equivalent:
|
||||
# mixes = F.linear(x_flat, fn)
|
||||
#
|
||||
# Scalar form:
|
||||
# mixes[token, o] +=
|
||||
# pre_fn[o, route, hidden] * cur_residual[route, hidden]
|
||||
#
|
||||
# This CTA computes only tile_mix_outputs columns of mixes.
|
||||
for tile_col_idx in T.unroll(tile_mix_outputs):
|
||||
mix_output_idx = (
|
||||
mix_output_tile_idx * tile_mix_outputs + tile_col_idx
|
||||
)
|
||||
|
||||
if mix_output_idx < num_mix_outputs:
|
||||
for route_idx in T.unroll(hc):
|
||||
mix_acc[tile_col_idx] += (
|
||||
pre_fn[mix_output_idx, route_idx, hidden_idx]
|
||||
* cur_residual_values[route_idx]
|
||||
)
|
||||
|
||||
# Reduce thread partials inside each warp.
|
||||
for tile_col_idx in T.unroll(tile_mix_outputs):
|
||||
mix_acc[tile_col_idx] = T.warp_reduce_sum(mix_acc[tile_col_idx])
|
||||
|
||||
if mix_output_tile_idx == 0:
|
||||
sqrsum_acc[0] = T.warp_reduce_sum(sqrsum_acc[0])
|
||||
|
||||
# One lane per warp writes warp-level partials to shared memory.
|
||||
if lane_idx == 0:
|
||||
for tile_col_idx in T.unroll(tile_mix_outputs):
|
||||
warp_partials[warp_idx, tile_col_idx] = mix_acc[tile_col_idx]
|
||||
|
||||
if mix_output_tile_idx == 0:
|
||||
warp_partials[warp_idx, tile_mix_outputs] = sqrsum_acc[0]
|
||||
|
||||
T.sync_threads()
|
||||
|
||||
# Reduce across warps and write split partials.
|
||||
#
|
||||
# The full PyTorch result would be:
|
||||
# mixes = F.linear(cur_residual.reshape(T, HC * H), fn)
|
||||
# sqrsum = (cur_residual.float() ** 2).sum(dim=(1, 2))
|
||||
#
|
||||
# This kernel is split along hidden, so each CTA writes only:
|
||||
# mixes_partial_out[hidden_split_idx, token, o]
|
||||
# sqrsum_partial_out[hidden_split_idx, token]
|
||||
#
|
||||
# Later mhc_pre_big_fuse does:
|
||||
# mixes = mixes_partial_out.sum(dim=0)
|
||||
# sqrsum = sqrsum_partial_out.sum(dim=0)
|
||||
# rms = rsqrt(sqrsum / (HC * H) + eps)
|
||||
# mixes *= rms
|
||||
# mixes -> pre/post/comb
|
||||
# layer_input = sum_j pre[j] * cur_residual[j]
|
||||
if warp_idx == 0:
|
||||
for tile_col_idx in T.unroll(tile_mix_outputs):
|
||||
mix_output_idx = mix_output_tile_idx * tile_mix_outputs + tile_col_idx
|
||||
|
||||
if mix_output_idx < num_mix_outputs and lane_idx == tile_col_idx:
|
||||
mix_output_partial = T.alloc_var(T.float32, init=0.0)
|
||||
|
||||
for reduce_warp_idx in T.unroll(num_warps):
|
||||
mix_output_partial += warp_partials[
|
||||
reduce_warp_idx, tile_col_idx
|
||||
]
|
||||
|
||||
mixes_partial_out[hidden_split_idx, token_idx, mix_output_idx] = (
|
||||
mix_output_partial
|
||||
)
|
||||
|
||||
if mix_output_tile_idx == 0 and lane_idx == 0:
|
||||
sqrsum_partial = T.alloc_var(T.float32, init=0.0)
|
||||
|
||||
for reduce_warp_idx in T.unroll(num_warps):
|
||||
sqrsum_partial += warp_partials[reduce_warp_idx, tile_mix_outputs]
|
||||
|
||||
sqrsum_partial_out[hidden_split_idx, token_idx] = sqrsum_partial
|
||||
|
||||
if ENABLE_PDL:
|
||||
T.pdl_trigger()
|
||||
|
||||
|
||||
def mhc_fused_post_pre(
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor,
|
||||
post_layer_mix: torch.Tensor,
|
||||
comb_res_mix: torch.Tensor,
|
||||
fn: torch.Tensor,
|
||||
hc_scale: torch.Tensor,
|
||||
hc_base: torch.Tensor,
|
||||
rms_eps: float,
|
||||
hc_pre_eps: float,
|
||||
hc_sinkhorn_eps: float,
|
||||
hc_post_mult_value: float,
|
||||
sinkhorn_repeat: int,
|
||||
n_splits: int = 1,
|
||||
tile_n: int = 1,
|
||||
*,
|
||||
norm_weight: torch.Tensor | None = None,
|
||||
norm_eps: float | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Fuse the boundary between one mHC post step and the next mHC pre step.
|
||||
|
||||
The unfused sequence is ``mhc_post -> pre-norm GEMM -> mhc_pre big_fuse``.
|
||||
This wrapper keeps the numerically sensitive ``mhc_pre_big_fuse`` stage,
|
||||
including optional RMSNorm, but removes the separate post/pre boundary.
|
||||
Small token batches use the FMA kernel above to combine ``mhc_post`` and the
|
||||
pre-norm GEMM in one launch; larger batches keep DeepGEMM for throughput and
|
||||
only fuse the Python/model-level scheduling boundary.
|
||||
|
||||
Returns:
|
||||
residual_cur: post-mapped residual, shape (..., hc_mult, hidden_size)
|
||||
post_mix_cur: shape (..., hc_mult, 1)
|
||||
comb_mix_cur: shape (..., hc_mult, hc_mult)
|
||||
layer_input_cur: shape (..., hidden_size)
|
||||
"""
|
||||
|
||||
assert residual.dtype == torch.bfloat16
|
||||
assert x.dtype == torch.bfloat16
|
||||
assert post_layer_mix.dtype == torch.float32
|
||||
assert comb_res_mix.dtype == torch.float32
|
||||
assert fn.dtype == torch.float32
|
||||
assert hc_scale.dtype == torch.float32
|
||||
assert hc_base.dtype == torch.float32
|
||||
|
||||
hc_mult = residual.shape[-2]
|
||||
hidden_size = residual.shape[-1]
|
||||
hc_mult2 = hc_mult * hc_mult
|
||||
hc_mult3 = hc_mult * 2 + hc_mult2
|
||||
hc_hidden_size = hc_mult * hidden_size
|
||||
outer_shape = residual.shape[:-2]
|
||||
|
||||
assert x.shape == (*outer_shape, hidden_size)
|
||||
assert post_layer_mix.shape in (
|
||||
(*outer_shape, hc_mult, 1),
|
||||
(*outer_shape, hc_mult),
|
||||
)
|
||||
assert comb_res_mix.shape == (*outer_shape, hc_mult, hc_mult)
|
||||
assert fn.shape == (hc_mult3, hc_hidden_size)
|
||||
assert hc_scale.shape == (3,)
|
||||
assert hc_base.shape == (hc_mult3,)
|
||||
|
||||
residual_flat = residual.view(-1, hc_mult, hidden_size)
|
||||
num_tokens = residual_flat.shape[0]
|
||||
if num_tokens == 0:
|
||||
# Some DP/EP ranks can receive no tokens; return correctly typed empty
|
||||
# tensors so later fused layers keep the same contracts as mhc_pre/hc_post.
|
||||
return (
|
||||
torch.empty_like(residual),
|
||||
torch.empty(
|
||||
(*outer_shape, hc_mult, 1), dtype=torch.float32, device=residual.device
|
||||
),
|
||||
torch.empty(
|
||||
(*outer_shape, hc_mult, hc_mult),
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
),
|
||||
torch.empty(
|
||||
(*outer_shape, hidden_size),
|
||||
dtype=torch.bfloat16,
|
||||
device=residual.device,
|
||||
),
|
||||
)
|
||||
x_flat = x.view(num_tokens, hidden_size)
|
||||
|
||||
# The scalar-FMA kernel wins only for small batches where launch
|
||||
# overhead dominates; beyond the threshold DeepGEMM's tensor-core path wins.
|
||||
fma_token_threshold = 32
|
||||
if num_tokens <= fma_token_threshold:
|
||||
tile_n = 2 if num_tokens < 8 else 3
|
||||
n_splits = 8 if (num_tokens < 8 and hidden_size <= 4096) else 4
|
||||
else:
|
||||
n_splits = _compute_num_split_for_mhc_pre(num_tokens, hc_hidden_size)
|
||||
|
||||
gemm_out_mul = torch.empty(
|
||||
n_splits,
|
||||
num_tokens,
|
||||
hc_mult3,
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
)
|
||||
gemm_out_sqrsum = torch.empty(
|
||||
n_splits,
|
||||
num_tokens,
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
)
|
||||
residual_cur = torch.empty_like(residual_flat)
|
||||
|
||||
if num_tokens <= fma_token_threshold:
|
||||
# Small-batch path: one TileLang launch computes hc_post, the bf16
|
||||
# residual write, GEMM partials, and the RMS square-sum partials.
|
||||
mhc_fused_post_pre_fma_tilelang(
|
||||
comb_res_mix.view(num_tokens, hc_mult, hc_mult),
|
||||
residual_flat,
|
||||
post_layer_mix.view(num_tokens, hc_mult),
|
||||
x_flat,
|
||||
fn.view(hc_mult3, hc_mult, hidden_size),
|
||||
gemm_out_mul,
|
||||
gemm_out_sqrsum,
|
||||
residual_cur,
|
||||
hc_mult,
|
||||
hidden_size,
|
||||
hc_mult3,
|
||||
tile_mix_outputs=tile_n,
|
||||
split_k=n_splits,
|
||||
)
|
||||
else:
|
||||
# Large-batch path: keep the existing high-throughput TileLang hc_post +
|
||||
# DeepGEMM pre-norm GEMM decomposition instead of replacing tensor cores.
|
||||
mhc_post_tilelang(
|
||||
comb_res_mix.view(num_tokens, hc_mult, hc_mult),
|
||||
residual_flat,
|
||||
post_layer_mix.view(num_tokens, hc_mult),
|
||||
x_flat,
|
||||
residual_cur,
|
||||
hc_mult,
|
||||
hidden_size,
|
||||
)
|
||||
|
||||
if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get():
|
||||
import deep_gemm
|
||||
|
||||
deep_gemm.tf32_hc_prenorm_gemm(
|
||||
residual_cur.view(num_tokens, hc_hidden_size),
|
||||
fn,
|
||||
gemm_out_mul,
|
||||
gemm_out_sqrsum,
|
||||
num_splits=n_splits,
|
||||
)
|
||||
else:
|
||||
# Fallback mirrors mhc_pre when DeepGEMM prenorm is disabled.
|
||||
n_splits = 1
|
||||
gemm_out_mul_2d = torch.empty(
|
||||
num_tokens, hc_mult3, dtype=torch.float32, device=residual.device
|
||||
)
|
||||
gemm_out_sqrsum_1d = torch.empty(
|
||||
num_tokens, dtype=torch.float32, device=residual.device
|
||||
)
|
||||
mhc_pre_gemm_sqrsum_tilelang(
|
||||
residual_cur.view(num_tokens, hc_hidden_size),
|
||||
fn,
|
||||
gemm_out_mul_2d,
|
||||
gemm_out_sqrsum_1d,
|
||||
hc_mult3,
|
||||
hc_hidden_size,
|
||||
)
|
||||
gemm_out_mul = gemm_out_mul_2d.unsqueeze(0)
|
||||
gemm_out_sqrsum = gemm_out_sqrsum_1d.unsqueeze(0)
|
||||
|
||||
post_mix_cur = torch.empty(
|
||||
num_tokens,
|
||||
hc_mult,
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
)
|
||||
comb_mix_cur = torch.empty(
|
||||
num_tokens,
|
||||
hc_mult2,
|
||||
dtype=torch.float32,
|
||||
device=residual.device,
|
||||
)
|
||||
layer_input_cur = torch.empty(
|
||||
num_tokens,
|
||||
hidden_size,
|
||||
dtype=torch.bfloat16,
|
||||
device=residual.device,
|
||||
)
|
||||
|
||||
if norm_weight is not None:
|
||||
# Final mhc_pre stage: convert GEMM partials into post/comb/layer_input
|
||||
# and fuse the following RMSNorm when the model passed a norm weight.
|
||||
assert norm_eps is not None
|
||||
assert norm_weight.shape == (hidden_size,)
|
||||
norm_weight_bf = (
|
||||
norm_weight.bfloat16()
|
||||
if norm_weight.dtype != torch.bfloat16
|
||||
else norm_weight
|
||||
)
|
||||
if not norm_weight_bf.is_contiguous():
|
||||
norm_weight_bf = norm_weight_bf.contiguous()
|
||||
mhc_pre_big_fuse_with_norm_tilelang(
|
||||
gemm_out_mul,
|
||||
gemm_out_sqrsum,
|
||||
hc_scale,
|
||||
hc_base,
|
||||
residual_cur,
|
||||
post_mix_cur,
|
||||
comb_mix_cur,
|
||||
layer_input_cur,
|
||||
norm_weight_bf,
|
||||
hidden_size,
|
||||
rms_eps,
|
||||
hc_pre_eps,
|
||||
hc_sinkhorn_eps,
|
||||
hc_post_mult_value,
|
||||
sinkhorn_repeat,
|
||||
norm_eps,
|
||||
n_splits,
|
||||
hc_mult,
|
||||
hc_mult3,
|
||||
)
|
||||
else:
|
||||
# Same mhc_pre finalization without the model-layer RMSNorm.
|
||||
mhc_pre_big_fuse_tilelang(
|
||||
gemm_out_mul,
|
||||
gemm_out_sqrsum,
|
||||
hc_scale,
|
||||
hc_base,
|
||||
residual_cur,
|
||||
post_mix_cur,
|
||||
comb_mix_cur,
|
||||
layer_input_cur,
|
||||
hidden_size,
|
||||
rms_eps,
|
||||
hc_pre_eps,
|
||||
hc_sinkhorn_eps,
|
||||
hc_post_mult_value,
|
||||
sinkhorn_repeat,
|
||||
n_splits,
|
||||
hc_mult,
|
||||
hc_mult3,
|
||||
)
|
||||
|
||||
return (
|
||||
residual_cur.view(*outer_shape, hc_mult, hidden_size),
|
||||
post_mix_cur.view(*outer_shape, hc_mult, 1),
|
||||
comb_mix_cur.view(*outer_shape, hc_mult, hc_mult),
|
||||
layer_input_cur.view(*outer_shape, hidden_size),
|
||||
)
|
||||
|
||||
@@ -2,6 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import concurrent.futures
|
||||
import logging
|
||||
import time
|
||||
from contextlib import nullcontext
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
@@ -61,6 +62,7 @@ from sglang.srt.layers.dp_attention import (
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||
from sglang.srt.layers.mhc import mhc_fused_post_pre
|
||||
from sglang.srt.layers.moe import get_moe_a2a_backend
|
||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
||||
from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8
|
||||
@@ -110,6 +112,18 @@ from sglang.srt.utils.hf_transformers_utils import get_rope_config
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_FP8_WO_A_GEMM = envs.SGLANG_OPT_FP8_WO_A_GEMM.get()
|
||||
_MHC_POST_MULT_VALUE = 2.0
|
||||
|
||||
|
||||
def _is_fused_mhc_post_pre_enabled() -> bool:
|
||||
# The fused path directly reuses TileLang mhc_post/mhc_pre kernels and their
|
||||
# tensor layout assumptions, so keep it disabled when either dependency is off.
|
||||
return (
|
||||
envs.SGLANG_OPT_FUSE_MHC_POST_PRE.get()
|
||||
and envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get()
|
||||
and envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get()
|
||||
)
|
||||
|
||||
|
||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||
_is_gfx95_supported = is_gfx95_supported()
|
||||
@@ -976,6 +990,133 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
self.hc_ffn_scale = nn.Parameter(torch.empty(3, dtype=torch.float32))
|
||||
self.rms_norm_eps = config.rms_norm_eps
|
||||
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
|
||||
self.use_fused_mhc_post_pre = _is_fused_mhc_post_pre_enabled()
|
||||
self._input_layernorm_weight_bf16 = None
|
||||
self._post_attention_layernorm_weight_bf16 = None
|
||||
|
||||
def refresh_mhc_norm_weight_cache(self):
|
||||
# Cache bf16 norm weights so the fused path does not allocate/cast per forward.
|
||||
self._input_layernorm_weight_bf16 = (
|
||||
self.input_layernorm.weight.data.bfloat16().contiguous()
|
||||
)
|
||||
self._post_attention_layernorm_weight_bf16 = (
|
||||
self.post_attention_layernorm.weight.data.bfloat16().contiguous()
|
||||
)
|
||||
|
||||
def prewarm_mhc_token_counts(
|
||||
self, token_counts: Tuple[int, ...], device: torch.device
|
||||
) -> None:
|
||||
paths = (
|
||||
(
|
||||
"attn",
|
||||
self.hc_attn_fn,
|
||||
self.hc_attn_scale,
|
||||
self.hc_attn_base,
|
||||
self.input_layernorm,
|
||||
),
|
||||
(
|
||||
"ffn",
|
||||
self.hc_ffn_fn,
|
||||
self.hc_ffn_scale,
|
||||
self.hc_ffn_base,
|
||||
self.post_attention_layernorm,
|
||||
),
|
||||
)
|
||||
|
||||
with torch.inference_mode():
|
||||
for num_tokens in token_counts:
|
||||
for path_name, hc_fn, hc_scale, hc_base, norm in paths:
|
||||
tic = time.perf_counter()
|
||||
residual = torch.empty(
|
||||
(num_tokens, self.hc_mult, self.hidden_size),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
y, post, comb, _ = self.hc_pre(
|
||||
residual,
|
||||
hc_fn,
|
||||
hc_scale,
|
||||
hc_base,
|
||||
norm=norm,
|
||||
)
|
||||
del residual, y, post, comb
|
||||
torch.cuda.synchronize()
|
||||
logger.info(
|
||||
"DeepSeek V4 MHC prewarm path=%s num_tokens=%s completed in %.3fs",
|
||||
path_name,
|
||||
num_tokens,
|
||||
time.perf_counter() - tic,
|
||||
)
|
||||
|
||||
if self.use_fused_mhc_post_pre:
|
||||
for num_tokens in token_counts:
|
||||
for path_name, hc_fn, hc_scale, hc_base, norm in paths:
|
||||
tic = time.perf_counter()
|
||||
# Dummy inputs matching the fused kernel's expected shapes.
|
||||
x = torch.empty(
|
||||
(num_tokens, self.hidden_size),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
residual = torch.empty(
|
||||
(num_tokens, self.hc_mult, self.hidden_size),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
post_mix = torch.empty(
|
||||
(num_tokens, self.hc_mult, 1),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
comb_mix = torch.empty(
|
||||
(num_tokens, self.hc_mult, self.hc_mult),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
norm_weight = norm.weight.data.bfloat16().contiguous()
|
||||
mhc_fused_post_pre(
|
||||
x,
|
||||
residual,
|
||||
post_mix,
|
||||
comb_mix,
|
||||
hc_fn,
|
||||
hc_scale,
|
||||
hc_base,
|
||||
self.rms_norm_eps,
|
||||
self.hc_eps,
|
||||
self.hc_eps,
|
||||
_MHC_POST_MULT_VALUE,
|
||||
self.hc_sinkhorn_iters,
|
||||
norm_weight=norm_weight,
|
||||
norm_eps=norm.variance_epsilon,
|
||||
)
|
||||
del x, residual, post_mix, comb_mix, norm_weight
|
||||
torch.cuda.synchronize()
|
||||
logger.info(
|
||||
"DeepSeek V4 MHC fused prewarm path=%s num_tokens=%s completed in %.3fs",
|
||||
path_name,
|
||||
num_tokens,
|
||||
time.perf_counter() - tic,
|
||||
)
|
||||
|
||||
def prewarm_mhc_token_count_buckets(
|
||||
self, max_num_tokens: int, device: torch.device
|
||||
) -> Tuple[int, ...]:
|
||||
from sglang.srt.layers.mhc import get_mhc_pre_token_count_representatives
|
||||
|
||||
token_counts = get_mhc_pre_token_count_representatives(
|
||||
max_num_tokens, self.hc_mult * self.hidden_size
|
||||
)
|
||||
if not token_counts:
|
||||
return token_counts
|
||||
|
||||
logger.info(
|
||||
"DeepSeek V4 MHC prewarm max_num_tokens=%s representative token counts: %s",
|
||||
max_num_tokens,
|
||||
token_counts,
|
||||
)
|
||||
self.prewarm_mhc_token_counts(token_counts, device)
|
||||
return token_counts
|
||||
|
||||
def hc_pre(
|
||||
self,
|
||||
@@ -1001,9 +1142,9 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
|
||||
if x.shape[0] == 0:
|
||||
y = torch.empty((0, shape[-1]), dtype=dtype, device=x.device)
|
||||
post = torch.empty((0, self.hc_mult), dtype=dtype, device=x.device)
|
||||
post = torch.empty((0, self.hc_mult), dtype=torch.float32, device=x.device)
|
||||
comb = torch.empty(
|
||||
(0, self.hc_mult, self.hc_mult), dtype=dtype, device=x.device
|
||||
(0, self.hc_mult, self.hc_mult), dtype=torch.float32, device=x.device
|
||||
)
|
||||
return y, post, comb, False
|
||||
|
||||
@@ -1023,7 +1164,7 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
rms_eps=self.rms_norm_eps,
|
||||
hc_pre_eps=self.hc_eps,
|
||||
hc_sinkhorn_eps=self.hc_eps,
|
||||
hc_post_mult_value=2.0,
|
||||
hc_post_mult_value=_MHC_POST_MULT_VALUE,
|
||||
sinkhorn_repeat=self.hc_sinkhorn_iters,
|
||||
**norm_kwargs,
|
||||
)
|
||||
@@ -1040,7 +1181,7 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
rms_eps=self.rms_norm_eps,
|
||||
hc_pre_eps=self.hc_eps,
|
||||
hc_sinkhorn_eps=self.hc_eps,
|
||||
hc_post_mult_value=2.0,
|
||||
hc_post_mult_value=_MHC_POST_MULT_VALUE,
|
||||
sinkhorn_repeat=self.hc_sinkhorn_iters,
|
||||
)
|
||||
return y, post.squeeze(-1), comb, False
|
||||
@@ -1122,27 +1263,60 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
input_ids: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
input_ids_global: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
residual = hidden_states
|
||||
hidden_states, post, comb, norm_fused = self.hc_pre(
|
||||
hidden_states,
|
||||
self.hc_attn_fn,
|
||||
self.hc_attn_scale,
|
||||
self.hc_attn_base,
|
||||
norm=self.input_layernorm,
|
||||
)
|
||||
if not norm_fused:
|
||||
if _use_aiter and _is_gfx95_supported:
|
||||
x_quant, hidden_states = _fused_rmsnorm_fp8_quant(
|
||||
hidden_states,
|
||||
self.input_layernorm.weight,
|
||||
self.rms_norm_eps,
|
||||
)
|
||||
else:
|
||||
hidden_states = self.input_layernorm(hidden_states)
|
||||
x_quant = None
|
||||
else:
|
||||
prev_residual: Optional[torch.Tensor] = None,
|
||||
prev_post: Optional[torch.Tensor] = None,
|
||||
prev_comb: Optional[torch.Tensor] = None,
|
||||
) -> Tuple[
|
||||
torch.Tensor,
|
||||
Optional[torch.Tensor],
|
||||
Optional[torch.Tensor],
|
||||
Optional[torch.Tensor],
|
||||
]:
|
||||
use_fused = self.use_fused_mhc_post_pre
|
||||
|
||||
if prev_residual is not None and use_fused:
|
||||
residual, post, comb, hidden_states = mhc_fused_post_pre(
|
||||
hidden_states,
|
||||
prev_residual,
|
||||
prev_post,
|
||||
prev_comb,
|
||||
self.hc_attn_fn,
|
||||
self.hc_attn_scale,
|
||||
self.hc_attn_base,
|
||||
self.rms_norm_eps,
|
||||
self.hc_eps,
|
||||
self.hc_eps,
|
||||
_MHC_POST_MULT_VALUE,
|
||||
self.hc_sinkhorn_iters,
|
||||
norm_weight=(
|
||||
self._input_layernorm_weight_bf16
|
||||
if self._input_layernorm_weight_bf16 is not None
|
||||
else self.input_layernorm.weight.data
|
||||
),
|
||||
norm_eps=self.input_layernorm.variance_epsilon,
|
||||
)
|
||||
x_quant = None
|
||||
else:
|
||||
residual = hidden_states
|
||||
hidden_states, post, comb, norm_fused = self.hc_pre(
|
||||
hidden_states,
|
||||
self.hc_attn_fn,
|
||||
self.hc_attn_scale,
|
||||
self.hc_attn_base,
|
||||
norm=self.input_layernorm,
|
||||
)
|
||||
if not norm_fused:
|
||||
if _use_aiter and _is_gfx95_supported:
|
||||
x_quant, hidden_states = _fused_rmsnorm_fp8_quant(
|
||||
hidden_states,
|
||||
self.input_layernorm.weight,
|
||||
self.rms_norm_eps,
|
||||
)
|
||||
else:
|
||||
hidden_states = self.input_layernorm(hidden_states)
|
||||
x_quant = None
|
||||
else:
|
||||
x_quant = None
|
||||
|
||||
hidden_states = self.self_attn(
|
||||
x=hidden_states,
|
||||
@@ -1151,35 +1325,58 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
x_quant=x_quant,
|
||||
)
|
||||
|
||||
fused_mhc = try_fused_hc_post_pre(
|
||||
hidden_states,
|
||||
residual,
|
||||
post,
|
||||
comb,
|
||||
self.hc_ffn_fn.T,
|
||||
self.hc_ffn_scale,
|
||||
self.hc_ffn_base,
|
||||
self.hc_mult,
|
||||
self.rms_norm_eps,
|
||||
self.hc_eps,
|
||||
2.0,
|
||||
self.hc_sinkhorn_iters,
|
||||
_is_gfx95_supported,
|
||||
)
|
||||
if fused_mhc is not None:
|
||||
residual, hidden_states, post, comb, norm_fused = fused_mhc
|
||||
if use_fused:
|
||||
fused_mhc = try_fused_hc_post_pre(
|
||||
hidden_states,
|
||||
residual,
|
||||
post,
|
||||
comb,
|
||||
self.hc_ffn_fn.T,
|
||||
self.hc_ffn_scale,
|
||||
self.hc_ffn_base,
|
||||
self.hc_mult,
|
||||
self.rms_norm_eps,
|
||||
self.hc_eps,
|
||||
_MHC_POST_MULT_VALUE,
|
||||
self.hc_sinkhorn_iters,
|
||||
_is_gfx95_supported,
|
||||
)
|
||||
if fused_mhc is not None:
|
||||
residual, hidden_states, post, comb, norm_fused = fused_mhc
|
||||
else:
|
||||
residual, post, comb, hidden_states = mhc_fused_post_pre(
|
||||
hidden_states,
|
||||
residual,
|
||||
post.unsqueeze(-1) if post.ndim == 2 else post,
|
||||
comb,
|
||||
self.hc_ffn_fn,
|
||||
self.hc_ffn_scale,
|
||||
self.hc_ffn_base,
|
||||
self.rms_norm_eps,
|
||||
self.hc_eps,
|
||||
self.hc_eps,
|
||||
_MHC_POST_MULT_VALUE,
|
||||
self.hc_sinkhorn_iters,
|
||||
norm_weight=(
|
||||
self._post_attention_layernorm_weight_bf16
|
||||
if self._post_attention_layernorm_weight_bf16 is not None
|
||||
else self.post_attention_layernorm.weight.data
|
||||
),
|
||||
norm_eps=self.post_attention_layernorm.variance_epsilon,
|
||||
)
|
||||
norm_fused = True
|
||||
else:
|
||||
hidden_states = self.hc_post(hidden_states, residual, post, comb)
|
||||
residual = hidden_states # [n, hc, d]
|
||||
residual = hidden_states
|
||||
hidden_states, post, comb, norm_fused = self.hc_pre(
|
||||
hidden_states,
|
||||
self.hc_ffn_fn,
|
||||
self.hc_ffn_scale,
|
||||
self.hc_ffn_base,
|
||||
norm=self.post_attention_layernorm,
|
||||
) # -> [n, d]
|
||||
if not norm_fused:
|
||||
hidden_states = self.post_attention_layernorm(hidden_states)
|
||||
)
|
||||
if not norm_fused:
|
||||
hidden_states = self.post_attention_layernorm(hidden_states)
|
||||
|
||||
_use_cp = self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch)
|
||||
_use_tp_moe_gather = (
|
||||
@@ -1233,9 +1430,13 @@ class DeepseekV4DecoderLayer(nn.Module):
|
||||
attn_tp_all_gather(gathered, hidden_states.contiguous())
|
||||
hidden_states = torch.cat(gathered)
|
||||
|
||||
hidden_states = self.hc_post(hidden_states, residual, post, comb)
|
||||
if not use_fused:
|
||||
hidden_states = self.hc_post(hidden_states, residual, post, comb)
|
||||
return hidden_states, None, None, None
|
||||
|
||||
return hidden_states
|
||||
# Return the deferred FFN hc_post state; the next layer consumes it with
|
||||
# cross-layer fusion, and the final layer is completed in DeepseekV4Model.
|
||||
return hidden_states, residual, post, comb
|
||||
|
||||
|
||||
class DeepseekV4Model(nn.Module):
|
||||
@@ -1302,6 +1503,7 @@ class DeepseekV4Model(nn.Module):
|
||||
self.hc_head_scale = nn.Parameter(torch.empty(1, dtype=torch.float32))
|
||||
|
||||
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
|
||||
self.use_fused_mhc_post_pre = _is_fused_mhc_post_pre_enabled()
|
||||
if self.dsa_enable_prefill_cp:
|
||||
self.cp_size = get_attention_cp_size()
|
||||
|
||||
@@ -1375,21 +1577,32 @@ class DeepseekV4Model(nn.Module):
|
||||
# forks alt-streams; later per-layer calls become no-ops.
|
||||
get_attn_backend()._maybe_upgrade_forward_metadata()
|
||||
|
||||
use_fused = self.use_fused_mhc_post_pre
|
||||
prev_residual, prev_post, prev_comb = None, None, None
|
||||
last_layer = None
|
||||
for i in range(self.start_layer, self.end_layer):
|
||||
layer = self.layers[i]
|
||||
last_layer = layer
|
||||
ctx = (
|
||||
nullcontext()
|
||||
if not get_global_server_args().disable_piecewise_cuda_graph
|
||||
else get_global_expert_distribution_recorder().with_current_layer(i)
|
||||
)
|
||||
with ctx:
|
||||
hidden_states = layer(
|
||||
hidden_states, prev_residual, prev_post, prev_comb = layer(
|
||||
positions=positions,
|
||||
hidden_states=hidden_states,
|
||||
forward_batch=forward_batch,
|
||||
input_ids=input_ids,
|
||||
input_ids_global=input_ids_global,
|
||||
prev_residual=prev_residual,
|
||||
prev_post=prev_post,
|
||||
prev_comb=prev_comb,
|
||||
)
|
||||
if use_fused and last_layer is not None:
|
||||
hidden_states = last_layer.hc_post(
|
||||
hidden_states, prev_residual, prev_post, prev_comb
|
||||
)
|
||||
|
||||
# CP all-gather only on the last PP rank; PP IPC carries CP-split tensors.
|
||||
if self.pp_group.is_last_rank and dsa_use_prefill_cp(forward_batch):
|
||||
@@ -1589,6 +1802,7 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
and not self_attn.indexer.compressor.ape_converted
|
||||
):
|
||||
self_attn.indexer.compressor.apply_ape_hotfix()
|
||||
layer.refresh_mhc_norm_weight_cache()
|
||||
|
||||
@staticmethod
|
||||
def remap_weight_name_to_dpsk_hf_format(
|
||||
|
||||
@@ -170,13 +170,17 @@ class DeepseekV4ModelNextN(nn.Module):
|
||||
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
|
||||
positions = cp_split_and_rebuild_position(forward_batch, positions)
|
||||
|
||||
hidden_states = self.decoder(
|
||||
hidden_states, residual, post, comb = self.decoder(
|
||||
positions=positions,
|
||||
hidden_states=hidden_states,
|
||||
forward_batch=forward_batch,
|
||||
input_ids=input_ids,
|
||||
input_ids_global=input_ids_global,
|
||||
)
|
||||
if residual is not None:
|
||||
# NextN has a single decoder layer, so no later layer can consume a
|
||||
# deferred fused hc_post state.
|
||||
hidden_states = self.decoder.hc_post(hidden_states, residual, post, comb)
|
||||
|
||||
if dsa_use_prefill_cp(forward_batch):
|
||||
hidden_states = cp_all_gather_rerange_output(
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import sglang.srt.layers.mhc as mhc
|
||||
from sglang.srt.layers.mhc import mhc_fused_post_pre, mhc_post, mhc_pre
|
||||
|
||||
|
||||
@pytest.mark.parametrize("hidden_size", [4096, 7168])
|
||||
@pytest.mark.parametrize("num_tokens", [0, 1, 8, 17, 32, 64])
|
||||
@pytest.mark.parametrize("use_norm", [False, True])
|
||||
def test_mhc_fused_post_pre_matches_unfused(
|
||||
monkeypatch, hidden_size, num_tokens, use_norm
|
||||
):
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA is required for TileLang mHC kernels")
|
||||
|
||||
monkeypatch.setattr(mhc, "is_dsa_prefill_cp_round_robin_split", lambda: False)
|
||||
torch.manual_seed(0)
|
||||
device = torch.device("cuda")
|
||||
hc_mult = 4
|
||||
hc_mult3 = hc_mult * 2 + hc_mult * hc_mult
|
||||
hc_hidden_size = hc_mult * hidden_size
|
||||
|
||||
x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) * 0.1
|
||||
residual = (
|
||||
torch.randn(
|
||||
num_tokens, hc_mult, hidden_size, device=device, dtype=torch.bfloat16
|
||||
)
|
||||
* 0.1
|
||||
)
|
||||
post_prev = torch.rand(num_tokens, hc_mult, 1, device=device, dtype=torch.float32)
|
||||
comb_prev = (
|
||||
torch.rand(num_tokens, hc_mult, hc_mult, device=device, dtype=torch.float32)
|
||||
* 0.25
|
||||
)
|
||||
fn = (
|
||||
torch.randn(hc_mult3, hc_hidden_size, device=device, dtype=torch.float32) * 0.01
|
||||
)
|
||||
hc_scale = torch.tensor([0.5, 0.25, 0.25], device=device, dtype=torch.float32)
|
||||
hc_base = torch.zeros(hc_mult3, device=device, dtype=torch.float32)
|
||||
norm_weight = (
|
||||
torch.ones(hidden_size, device=device, dtype=torch.bfloat16)
|
||||
if use_norm
|
||||
else None
|
||||
)
|
||||
norm_eps = 1e-6 if use_norm else None
|
||||
|
||||
rms_eps = 1e-6
|
||||
hc_eps = 1e-6
|
||||
sinkhorn_repeat = 2
|
||||
|
||||
residual_ref = post_ref = comb_ref = layer_ref = None
|
||||
if num_tokens > 0:
|
||||
residual_ref = mhc_post(x, residual, post_prev, comb_prev)
|
||||
post_ref, comb_ref, layer_ref = mhc_pre(
|
||||
residual_ref,
|
||||
fn,
|
||||
hc_scale,
|
||||
hc_base,
|
||||
rms_eps,
|
||||
hc_eps,
|
||||
hc_eps,
|
||||
2.0,
|
||||
sinkhorn_repeat,
|
||||
norm_weight=norm_weight,
|
||||
norm_eps=norm_eps,
|
||||
)
|
||||
residual_out, post_out, comb_out, layer_out = mhc_fused_post_pre(
|
||||
x,
|
||||
residual,
|
||||
post_prev,
|
||||
comb_prev,
|
||||
fn,
|
||||
hc_scale,
|
||||
hc_base,
|
||||
rms_eps,
|
||||
hc_eps,
|
||||
hc_eps,
|
||||
2.0,
|
||||
sinkhorn_repeat,
|
||||
norm_weight=norm_weight,
|
||||
norm_eps=norm_eps,
|
||||
)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
if num_tokens == 0:
|
||||
assert residual_out.shape == residual.shape
|
||||
assert post_out.shape == (0, hc_mult, 1)
|
||||
assert comb_out.shape == (0, hc_mult, hc_mult)
|
||||
assert layer_out.shape == (0, hidden_size)
|
||||
assert residual_out.dtype == torch.bfloat16
|
||||
assert post_out.dtype == torch.float32
|
||||
assert comb_out.dtype == torch.float32
|
||||
assert layer_out.dtype == torch.bfloat16
|
||||
return
|
||||
|
||||
assert residual_ref is not None
|
||||
assert post_ref is not None
|
||||
assert comb_ref is not None
|
||||
assert layer_ref is not None
|
||||
assert residual_out.shape == residual_ref.shape
|
||||
assert post_out.shape == post_ref.shape
|
||||
assert comb_out.shape == comb_ref.shape
|
||||
assert layer_out.shape == layer_ref.shape
|
||||
|
||||
torch.testing.assert_close(residual_out, residual_ref, atol=0, rtol=0)
|
||||
torch.testing.assert_close(post_out, post_ref, atol=1e-3, rtol=1e-3)
|
||||
torch.testing.assert_close(comb_out, comb_ref, atol=1e-3, rtol=1e-3)
|
||||
layer_atol = 2e-2 if use_norm else 2e-3
|
||||
layer_rtol = 2e-2 if use_norm else 2e-3
|
||||
torch.testing.assert_close(layer_out, layer_ref, atol=layer_atol, rtol=layer_rtol)
|
||||
Reference in New Issue
Block a user