diff --git a/python/sglang/test/layer_ut_utils.py b/python/sglang/test/layer_ut_utils.py new file mode 100644 index 000000000..01a9f24fc --- /dev/null +++ b/python/sglang/test/layer_ut_utils.py @@ -0,0 +1,74 @@ +"""Shared fixture plumbing for layer-level backend parity UTs. + +Hand-written quantization references (oracle side) live in quant_ref_utils. +""" + +import os + +import torch + + +def init_single_process_dist(master_port: int = 29632, backend: str = "gloo"): + """world=1 dist + model-parallel groups; srt layers require them even + at tp=1.""" + os.environ.setdefault("MASTER_ADDR", "127.0.0.1") + os.environ.setdefault("MASTER_PORT", str(master_port)) + os.environ.setdefault("RANK", "0") + os.environ.setdefault("WORLD_SIZE", "1") + os.environ.setdefault("LOCAL_RANK", "0") + from sglang.srt.distributed.parallel_state import ( + init_distributed_environment, + initialize_model_parallel, + model_parallel_is_initialized, + ) + + if not torch.distributed.is_initialized(): + init_distributed_environment( + world_size=1, rank=0, local_rank=0, backend=backend + ) + if not model_parallel_is_initialized(): + # kwargs only: a positional backend would land in the + # attention_data_parallel_size slot and explode on int // str. + initialize_model_parallel( + tensor_model_parallel_size=1, + expert_model_parallel_size=1, + pipeline_model_parallel_size=1, + backend=backend, + ) + + +def make_tp1_column_parallel_linear( + quant_config, n: int, k: int, prefix: str = "model.layers.0.mlp.up_proj", **kwargs +): + from sglang.srt.layers.linear import ColumnParallelLinear + + return ColumnParallelLinear( + input_size=k, + output_size=n, + bias=False, + params_dtype=torch.bfloat16, + quant_config=quant_config, + prefix=prefix, + tp_rank=0, + tp_size=1, + **kwargs, + ).cuda() + + +def load_linear_weights(layer, shard_id=None, **named_weights): + """Feed checkpoint-format tensors through the real weight_loader.""" + for name, loaded in named_weights.items(): + if shard_id is None: + layer.weight_loader_v2(getattr(layer, name), loaded) + else: + layer.weight_loader_v2(getattr(layer, name), loaded, shard_id) + + +def assert_output_close(tc, out, ref, cos_threshold=0.99, rtol=None, atol=None): + tc.assertEqual(tuple(out.shape), tuple(ref.shape)) + cos = torch.nn.functional.cosine_similarity( + out.float().flatten(), ref.flatten(), dim=0 + ).item() + tc.assertGreater(cos, cos_threshold) + if rtol is not None: + torch.testing.assert_close(out.float(), ref, rtol=rtol, atol=atol) diff --git a/python/sglang/test/quant_ref_utils.py b/python/sglang/test/quant_ref_utils.py new file mode 100644 index 000000000..6f8a19aa4 --- /dev/null +++ b/python/sglang/test/quant_ref_utils.py @@ -0,0 +1,70 @@ +"""Hand-written quantization-format references for backend parity UTs. + +Deliberately independent of sglang.srt -- never replace with srt imports; +the tests use these to check srt. +""" + +import torch + +FLOAT8_E4M3_MAX = 448.0 +FLOAT4_E2M1_MAX = 6.0 + +kE2M1ToFloat = torch.tensor( + [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=torch.float32 +) + + +def convert_swizzled_to_linear(a_sf_swizzled: torch.Tensor, m, k, block_size=16): + m_tiles = (m + 128 - 1) // 128 + f = block_size * 4 + k_tiles = (k + f - 1) // f + tmp = torch.reshape(a_sf_swizzled, (1, m_tiles, k_tiles, 32, 4, 4)) + tmp = torch.permute(tmp, (0, 1, 4, 3, 2, 5)) + out = tmp.reshape(m_tiles * 128, k_tiles * f // block_size) + # Crop the K-tile padding too: k // block_size scale columns, not k. + return out[0:m, 0 : k // block_size] + + +def break_fp4_bytes(a, dtype=torch.float32): + assert a.dtype == torch.uint8 + m, n = a.shape + a_flat = a.flatten() + high = (a_flat & 0xF0) >> 4 + low = a_flat & 0x0F + combined = torch.stack((low, high), dim=1).flatten() + signs = (combined & 0x08).to(torch.bool) + abs_vals = (combined & 0x07).to(torch.long) + kE2M1 = kE2M1ToFloat.to(device=a.device) + values = kE2M1[abs_vals] * torch.where(signs, -1.0, 1.0) + return values.reshape(m, n * 2).to(dtype=dtype) + + +def dequantize_nvfp4_to_dtype( + tensor_fp4, tensor_sf, global_scale, dtype, block_size=16 +): + assert tensor_fp4.dtype == torch.uint8 + m, packed_k = tensor_fp4.shape + k = packed_k * 2 + tensor_f32 = break_fp4_bytes(tensor_fp4, torch.float32) + tensor_f32 = tensor_f32.reshape(m, k // block_size, block_size) + tensor_sf = tensor_sf.view(torch.float8_e4m3fn) + tensor_sf = convert_swizzled_to_linear(tensor_sf, m, k, block_size) + tensor_sf_dtype = tensor_sf.to(torch.float32) / global_scale + out = (tensor_f32 * tensor_sf_dtype.unsqueeze(-1)).reshape(m, k) + return out.to(dtype=dtype) + + +def quantize_nvfp4_shard(w: torch.Tensor, gs=None): + """NVFP4-quantize one checkpoint shard; returns (packed, linear sf, + global scale, fp32 dequant reference).""" + from flashinfer import fp4_quantize + + n, k = w.shape + if gs is None: + gs = FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX / w.abs().max().to(torch.float32) + w_q, w_sf_swizzled = fp4_quantize(w, gs) + sf_linear = convert_swizzled_to_linear( + w_sf_swizzled.view(torch.float8_e4m3fn), n, k, 16 + ) + w_dequant = dequantize_nvfp4_to_dtype(w_q, w_sf_swizzled, gs, torch.float32) + return w_q, sf_linear, gs, w_dequant diff --git a/test/registered/debug_utils/test_tensor_dump_forward_hook.py b/test/registered/debug_utils/test_tensor_dump_forward_hook.py index 00dee4adf..5074f6a92 100644 --- a/test/registered/debug_utils/test_tensor_dump_forward_hook.py +++ b/test/registered/debug_utils/test_tensor_dump_forward_hook.py @@ -6,17 +6,14 @@ from torch import nn from sglang.srt.debug_utils.tensor_dump_forward_hook import ( register_forward_hook_for_model, ) -from sglang.srt.distributed.parallel_state import ( - get_default_distributed_backend, - init_distributed_environment, - initialize_model_parallel, -) +from sglang.srt.distributed.parallel_state import get_default_distributed_backend from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import LinearBase from sglang.srt.models.qwen2 import Qwen2MLP from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler from sglang.srt.utils import add_prefix, get_device from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.layer_ut_utils import init_single_process_dist register_cuda_ci( est_time=9, @@ -80,15 +77,7 @@ def init_weights(module): def test_model_forward_dump(tmp_path): set_global_server_args_for_scheduler(ServerArgs(model_path="dummy")) device = get_device() - backend = get_default_distributed_backend(device) - init_distributed_environment( - backend=backend, - world_size=1, - rank=0, - local_rank=0, - distributed_init_method="tcp://127.0.0.1:2646", - ) - initialize_model_parallel() + init_single_process_dist(backend=get_default_distributed_backend(device)) model = MockCausalLM() model.apply(init_weights) model = model.to(device=device, dtype=torch.bfloat16) diff --git a/test/registered/kernels/ops/moe/test_fp4_moe.py b/test/registered/kernels/ops/moe/test_fp4_moe.py index 9fb940eba..8263cfa3b 100644 --- a/test/registered/kernels/ops/moe/test_fp4_moe.py +++ b/test/registered/kernels/ops/moe/test_fp4_moe.py @@ -10,6 +10,11 @@ from torch.nn import functional as F from sglang.srt.layers.moe.topk import TopKConfig, select_experts from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.quant_ref_utils import ( + FLOAT4_E2M1_MAX, + FLOAT8_E4M3_MAX, + dequantize_nvfp4_to_dtype, +) register_cuda_ci(est_time=300, suite="nightly-4-gpu-b200", nightly=True) @@ -19,66 +24,6 @@ if torch.cuda.get_device_capability() < (10, 0): allow_module_level=True, ) -kE2M1ToFloat = torch.tensor( - [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=torch.float32 -) - -FLOAT8_E4M3_MAX = 448.0 -FLOAT4_E2M1_MAX = 6.0 - - -def convert_swizzled_to_linear(a_sf_swizzled: torch.Tensor, m, k, block_size): - m_tiles = (m + 128 - 1) // 128 - f = block_size * 4 - k_tiles = (k + f - 1) // f - tmp = torch.reshape(a_sf_swizzled, (1, m_tiles, k_tiles, 32, 4, 4)) - tmp = torch.permute(tmp, (0, 1, 4, 3, 2, 5)) - out = tmp.reshape(m_tiles * 128, k_tiles * f // block_size) - return out[0:m, 0:k] - - -def dequantize_nvfp4_to_dtype( - tensor_fp4, tensor_sf, global_scale, dtype, device, block_size=16 -): - """Dequantize the fp4 tensor back to high precision.""" - # Two fp4 values are packed into one uint8. - assert tensor_fp4.dtype == torch.uint8 - m, packed_k = tensor_fp4.shape - k = packed_k * 2 - tensor_f32 = break_fp4_bytes(tensor_fp4, dtype) - tensor_f32 = tensor_f32.reshape(m, k // block_size, block_size) - tensor_sf = tensor_sf.view(torch.float8_e4m3fn) - tensor_sf = convert_swizzled_to_linear(tensor_sf, m, k, block_size) - tensor_sf_dtype = tensor_sf.to(torch.float32) / global_scale - - # scale the tensor - out = (tensor_f32 * tensor_sf_dtype.unsqueeze(-1)).reshape(m, k) - return out.to(dtype=dtype) - - -def break_fp4_bytes(a, dtype): - assert a.dtype == torch.uint8 - m, n = a.shape - - # Vectorized nibble processing - a_flat = a.flatten() - high = (a_flat & 0xF0) >> 4 # Upper nibbles - low = a_flat & 0x0F # Lower nibbles - - # Combine nibbles for batch processing - combined = torch.stack((low, high), dim=1).flatten() - - # Vectorized sign and magnitude extraction - signs = (combined & 0x08).to(torch.bool) # Sign bits - abs_vals = (combined & 0x07).to(torch.long) # Magnitude indices - - # Device-aware lookup and sign application - kE2M1 = kE2M1ToFloat.to(device=a.device) - values = kE2M1[abs_vals] * torch.where(signs, -1.0, 1.0) - - # Reshape to final form - return values.reshape(m, n * 2).to(dtype=dtype) - def compute_routing(router_logits: torch.Tensor, top_k: int): routing_weights = torch.softmax(router_logits, dim=1, dtype=torch.float) @@ -170,7 +115,6 @@ def torch_moe_nvfp4(a, w1, w2, topk, topk_weight, topk_ids): inter_blockscale, inter_gs, dtype=inter.dtype, - device=inter.device, block_size=16, ).cuda() out[mask] = inter @ w2[i].transpose(0, 1) @@ -319,7 +263,6 @@ def check_moe( a_scale_interleaved, a_global_scale, dtype=a.dtype, - device=a.device, block_size=quant_blocksize, ) @@ -332,7 +275,6 @@ def check_moe( w1_blockscale[idx], w1_gs[idx], dtype=w1.dtype, - device=w1.device, block_size=quant_blocksize, ) w2_d[idx] = dequantize_nvfp4_to_dtype( @@ -340,7 +282,6 @@ def check_moe( w2_blockscale[idx], w2_gs[idx], dtype=w2.dtype, - device=w2.device, block_size=quant_blocksize, ) diff --git a/test/registered/moe/test_cutedsl_moe.py b/test/registered/moe/test_cutedsl_moe.py index 30e7dcf3d..452d187df 100644 --- a/test/registered/moe/test_cutedsl_moe.py +++ b/test/registered/moe/test_cutedsl_moe.py @@ -8,6 +8,11 @@ from torch.nn import functional as F from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.moe.flashinfer_cutedsl_moe import flashinfer_cutedsl_moe_masked from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.quant_ref_utils import ( + FLOAT4_E2M1_MAX, + FLOAT8_E4M3_MAX, + dequantize_nvfp4_to_dtype, +) try: from flashinfer import CuteDslMoEWrapper @@ -21,66 +26,6 @@ register_cuda_ci(est_time=24, stage="extra-b", runner_config="4-gpu-b200") SKIP_TEST = torch.cuda.get_device_capability() < (10, 0) SKIP_REASON = "Nvfp4 Requires compute capability of 10 or above." -kE2M1ToFloat = torch.tensor( - [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=torch.float32 -) - -FLOAT8_E4M3_MAX = 448.0 -FLOAT4_E2M1_MAX = 6.0 - - -def convert_swizzled_to_linear(a_sf_swizzled: torch.Tensor, m, k, block_size): - m_tiles = (m + 128 - 1) // 128 - f = block_size * 4 - k_tiles = (k + f - 1) // f - tmp = torch.reshape(a_sf_swizzled, (1, m_tiles, k_tiles, 32, 4, 4)) - tmp = torch.permute(tmp, (0, 1, 4, 3, 2, 5)) - out = tmp.reshape(m_tiles * 128, k_tiles * f // block_size) - return out[0:m, 0:k] - - -def dequantize_nvfp4_to_dtype( - tensor_fp4, tensor_sf, global_scale, dtype, device, block_size=16 -): - """Dequantize the fp4 tensor back to high precision.""" - # Two fp4 values are packed into one uint8. - assert tensor_fp4.dtype == torch.uint8 - m, packed_k = tensor_fp4.shape - k = packed_k * 2 - tensor_f32 = break_fp4_bytes(tensor_fp4, dtype) - tensor_f32 = tensor_f32.reshape(m, k // block_size, block_size) - tensor_sf = tensor_sf.view(torch.float8_e4m3fn) - tensor_sf = convert_swizzled_to_linear(tensor_sf, m, k, block_size) - tensor_sf_dtype = tensor_sf.to(torch.float32) / global_scale - - # scale the tensor - out = (tensor_f32 * tensor_sf_dtype.unsqueeze(-1)).reshape(m, k) - return out.to(dtype=dtype) - - -def break_fp4_bytes(a, dtype): - assert a.dtype == torch.uint8 - m, n = a.shape - - # Vectorized nibble processing - a_flat = a.flatten() - high = (a_flat & 0xF0) >> 4 # Upper nibbles - low = a_flat & 0x0F # Lower nibbles - - # Combine nibbles for batch processing - combined = torch.stack((low, high), dim=1).flatten() - - # Vectorized sign and magnitude extraction - signs = (combined & 0x08).to(torch.bool) # Sign bits - abs_vals = (combined & 0x07).to(torch.long) # Magnitude indices - - # Device-aware lookup and sign application - kE2M1 = kE2M1ToFloat.to(device=a.device) - values = kE2M1[abs_vals] * torch.where(signs, -1.0, 1.0) - - # Reshape to final form - return values.reshape(m, n * 2).to(dtype=dtype) - def _interleave_w13_halves( x: torch.Tensor, group_size: int = 64, dim: int = -1 @@ -464,7 +409,6 @@ def torch_moe_nvfp4(a, w1, w2, topk, topk_weight, topk_ids): inter_blockscale, inter_gs, dtype=inter.dtype, - device=inter.device, block_size=16, ).cuda() out[mask] = inter @ w2[i].transpose(0, 1) @@ -916,7 +860,6 @@ class TestCuteDslV1(unittest.TestCase): a_scale_interleaved, a_global_scale, dtype=hidden_states.dtype, - device=hidden_states.device, block_size=16, ) w1_d = torch.empty( @@ -942,7 +885,6 @@ class TestCuteDslV1(unittest.TestCase): w1_blockscale_sliced, w1_global_scale[idx], dtype=w1.dtype, - device=w1.device, block_size=16, ) w2_d[idx] = dequantize_nvfp4_to_dtype( @@ -950,7 +892,6 @@ class TestCuteDslV1(unittest.TestCase): w2_blockscale_sliced, w2_global_scale[idx], dtype=w2.dtype, - device=w2.device, block_size=16, ) @@ -1093,7 +1034,6 @@ class TestCuteDslV1(unittest.TestCase): a_scale_interleaved, a_global_scale, dtype=hidden_states.dtype, - device=device, block_size=16, ) w1_d = torch.empty( @@ -1117,7 +1057,6 @@ class TestCuteDslV1(unittest.TestCase): w1_blockscale_sliced, w1_global_scale[idx], dtype=w1.dtype, - device=device, block_size=16, ) w2_d[idx] = dequantize_nvfp4_to_dtype( @@ -1125,7 +1064,6 @@ class TestCuteDslV1(unittest.TestCase): w2_blockscale_sliced, w2_global_scale[idx], dtype=w2.dtype, - device=device, block_size=16, ) @@ -1268,7 +1206,6 @@ class TestCuteDslV1(unittest.TestCase): a_scale, a_gs, dtype=torch.bfloat16, - device=device, block_size=16, ) w1_d = torch.empty( @@ -1287,7 +1224,6 @@ class TestCuteDslV1(unittest.TestCase): w1_blockscale_sliced, w1_gs[idx], dtype=w1.dtype, - device=device, block_size=16, ) w2_d[idx] = dequantize_nvfp4_to_dtype( @@ -1295,7 +1231,6 @@ class TestCuteDslV1(unittest.TestCase): w2_blockscale_sliced, w2_gs[idx], dtype=w2.dtype, - device=device, block_size=16, ) ref = torch_moe_nvfp4( diff --git a/test/registered/moe/test_hpc_ops_moe.py b/test/registered/moe/test_hpc_ops_moe.py index 0e30bcb45..483d3a978 100644 --- a/test/registered/moe/test_hpc_ops_moe.py +++ b/test/registered/moe/test_hpc_ops_moe.py @@ -6,16 +6,10 @@ FP8 quantized weights. Skipped when HPC-Ops (https://github.com/Tencent/hpc-ops) is not installed or the GPU is not SM90 (the kernels ship sm90a only). """ -import os import unittest import torch -from sglang.srt.distributed.parallel_state import ( - init_distributed_environment, - initialize_model_parallel, - model_parallel_is_initialized, -) from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig from sglang.srt.layers.moe.moe_runner.hpc_ops import ( HpcOpsMoeQuantInfo, @@ -27,6 +21,7 @@ from sglang.srt.layers.moe.token_dispatcher.standard import StandardDispatchOutp from sglang.srt.layers.moe.topk import StandardTopKOutput from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.layer_ut_utils import init_single_process_dist from sglang.test.test_utils import CustomTestCase register_cuda_ci(est_time=60, stage="base-b", runner_config="1-gpu-large") @@ -43,26 +38,11 @@ def _sm90() -> bool: def _ensure_dist_initialized() -> None: - """Single-rank gloo distributed + model-parallel groups (TP=1, EP=1). - - The triton fused_experts reference allocates its output under + """The triton fused_experts reference allocates its output under ``use_symmetric_memory(get_tp_group(), ...)``, which requires the TP group even when symmetric allocation is disabled. """ - os.environ.setdefault("MASTER_ADDR", "127.0.0.1") - os.environ.setdefault("MASTER_PORT", "29633") - os.environ.setdefault("RANK", "0") - os.environ.setdefault("WORLD_SIZE", "1") - os.environ.setdefault("LOCAL_RANK", "0") - if not torch.distributed.is_initialized(): - init_distributed_environment(world_size=1, rank=0, local_rank=0, backend="gloo") - if not model_parallel_is_initialized(): - initialize_model_parallel( - tensor_model_parallel_size=1, - expert_model_parallel_size=1, - pipeline_model_parallel_size=1, - backend="gloo", - ) + init_single_process_dist(master_port=29633) def _quant_blockwise(w: torch.Tensor, block: int = 128): diff --git a/test/registered/quant/test_gptqmodel_dynamic.py b/test/registered/quant/test_gptqmodel_dynamic.py index 47e12d0ad..7c435e609 100644 --- a/test/registered/quant/test_gptqmodel_dynamic.py +++ b/test/registered/quant/test_gptqmodel_dynamic.py @@ -21,28 +21,14 @@ def check_quant_method(model_path: str, use_marlin_kernel: bool): from sglang.srt.configs.device_config import DeviceConfig from sglang.srt.configs.load_config import LoadConfig from sglang.srt.configs.model_config import ModelConfig - from sglang.srt.distributed import ( - init_distributed_environment, - initialize_model_parallel, - ) from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state from sglang.srt.layers.quantization.utils import get_dynamic_override from sglang.srt.model_loader import get_model from sglang.srt.server_args import ServerArgs + from sglang.test.layer_ut_utils import init_single_process_dist - try: - init_distributed_environment( - backend="nccl", - world_size=1, - rank=0, - local_rank=0, - distributed_init_method="tcp://127.0.0.1:2646", - ) - initialize_model_parallel(tensor_model_parallel_size=1) - monkey_patch_vllm_parallel_state() - except AssertionError: - # ignore this error: tensor model parallel group is already initialized - pass + init_single_process_dist(backend="nccl") + monkey_patch_vllm_parallel_state() server_args = ServerArgs(model_path=model_path, dtype=torch.float16) set_global_server_args_for_scheduler(server_args) diff --git a/test/registered/unit/layers/quantization/test_fp8_blockwise_linear_backends.py b/test/registered/unit/layers/quantization/test_fp8_blockwise_linear_backends.py index 04f675bd1..ac1d628be 100644 --- a/test/registered/unit/layers/quantization/test_fp8_blockwise_linear_backends.py +++ b/test/registered/unit/layers/quantization/test_fp8_blockwise_linear_backends.py @@ -1,13 +1,8 @@ """Numerics for the FP8 dense-linear GEMM backends (--fp8-gemm-backend). -Runs the quant-method layer path (create_weights -> -process_weights_after_loading -> apply) against a dequantized-reference -matmul, covering the per-backend weight preparation (e.g. UE8M0 scale requant -for DeepGEMM, per-backend MXFP8 scale packing) and the GEMM dispatch. -Three formats: FP8 blockwise (Fp8LinearMethod), MXFP8 (Fp8LinearMethod with -use_mxfp8), and per-tensor FP8 (ModelOptFp8LinearMethod, auto dispatch). -The backend set adapts to the device SM version, so the same file covers -Hopper (SM90), B200-class (SM100/103), and consumer Blackwell (SM120). +Real layer path vs a dequantized-reference matmul, in three formats: FP8 +blockwise, MXFP8, and per-tensor FP8 (auto dispatch). Backend sets adapt to +the device SM, so one file covers SM90 / SM100 / SM120. """ import unittest @@ -16,14 +11,17 @@ from unittest import mock import torch from sglang.srt.layers.quantization import fp8_utils -from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8LinearMethod +from sglang.srt.layers.quantization.fp8 import Fp8Config from sglang.srt.layers.quantization.fp8_utils import Fp8GemmRunnerBackend -from sglang.srt.layers.quantization.modelopt_quant import ( - ModelOptFp8Config, - ModelOptFp8LinearMethod, -) +from sglang.srt.layers.quantization.modelopt_quant import ModelOptFp8Config from sglang.srt.utils import get_device_sm from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.layer_ut_utils import ( + assert_output_close, + init_single_process_dist, + load_linear_weights, + make_tp1_column_parallel_linear, +) from sglang.test.test_utils import CustomTestCase register_cuda_ci(est_time=120, stage="base-b", runner_config="4-gpu-b200") @@ -99,26 +97,17 @@ def _quantize_mxfp8(w: torch.Tensor, block: int = 32): return w_fp8.reshape(n, k), scale_e8m0, w_dequant -def _create_weights(method, n: int, k: int, device: str = "cuda"): - layer = torch.nn.Module() - kwargs = {} - if isinstance(method, Fp8LinearMethod): - # The shape check reads TP world size (needs distributed init); skip it here. - kwargs["skip_block_quant_check"] = True - method.create_weights( - layer, - input_size_per_partition=k, - output_partition_sizes=[n], - input_size=k, - output_size=n, - params_dtype=torch.bfloat16, - weight_loader=lambda *args, **kw: None, - **kwargs, +def _make_linear(quant_config, n: int, k: int): + return make_tp1_column_parallel_linear( + quant_config, n, k, skip_block_quant_check=True ) - return layer.to(device) class _LinearBackendCheck(CustomTestCase): + @classmethod + def setUpClass(cls): + init_single_process_dist() + def _check_backend(self, backend: str, allowed, shapes, build_layer): if backend not in allowed: self.skipTest(f"{backend} not in SM{get_device_sm()} backend set") @@ -130,21 +119,16 @@ class _LinearBackendCheck(CustomTestCase): "FP8_GEMM_RUNNER_BACKEND", Fp8GemmRunnerBackend(backend), ): - method, layer, w_dequant = build_layer(n, k) - method.process_weights_after_loading(layer) + layer, w_dequant = build_layer(n, k) + layer.quant_method.process_weights_after_loading(layer) x = torch.randn((m, k), device="cuda", dtype=torch.bfloat16) / 10 - out = method.apply(layer, x) + out, _ = layer(x) ref = x.float() @ w_dequant.T - self.assertEqual(out.shape, (m, n)) - cos = torch.nn.functional.cosine_similarity( - out.float().flatten(), ref.flatten(), dim=0 - ).item() - self.assertGreater(cos, 0.99) # atol covers single-element UE8M0 scale-rounding outliers # (deep_gemm); a wrong kernel/layout fails by orders more. - torch.testing.assert_close(out.float(), ref, rtol=5e-2, atol=1e-1) + assert_output_close(self, out, ref, rtol=5e-2, atol=1e-1) @unittest.skipIf(get_device_sm() < 90, "FP8 GEMM backends require SM90+") @@ -156,13 +140,11 @@ class TestFp8BlockwiseLinearBackends(_LinearBackendCheck): activation_scheme="dynamic", weight_block_size=[128, 128], ) - method = Fp8LinearMethod(quant_config) - layer = _create_weights(method, n, k) + layer = _make_linear(quant_config, n, k) w = torch.randn((n, k), device="cuda", dtype=torch.bfloat16) / 10 w_fp8, scale_inv, w_dequant = _quantize_fp8_blockwise(w) - layer.weight.data.copy_(w_fp8) - layer.weight_scale_inv.data.copy_(scale_inv) - return method, layer, w_dequant + load_linear_weights(layer, weight=w_fp8, weight_scale_inv=scale_inv) + return layer, w_dequant def _run(self, backend: str): self._check_backend( @@ -197,13 +179,11 @@ class TestMxfp8LinearBackends(_LinearBackendCheck): activation_scheme="dynamic", use_mxfp8=True, ) - method = Fp8LinearMethod(quant_config) - layer = _create_weights(method, n, k) + layer = _make_linear(quant_config, n, k) w = torch.randn((n, k), device="cuda", dtype=torch.bfloat16) / 10 w_fp8, scale_e8m0, w_dequant = _quantize_mxfp8(w) - layer.weight.data.copy_(w_fp8) - layer.weight_scale_inv.data.copy_(scale_e8m0) - return method, layer, w_dequant + load_linear_weights(layer, weight=w_fp8, weight_scale_inv=scale_e8m0) + return layer, w_dequant def _run(self, backend: str): self._check_backend(backend, _mxfp8_backends(), MXFP8_SHAPES, self._build_layer) @@ -225,17 +205,22 @@ class TestModeloptFp8PerTensorLinear(_LinearBackendCheck): @staticmethod def _build_layer(n: int, k: int): - quant_config = ModelOptFp8Config(is_checkpoint_fp8_serialized=True) - method = ModelOptFp8LinearMethod(quant_config) - layer = _create_weights(method, n, k) + quant_config = ModelOptFp8Config( + is_checkpoint_fp8_serialized=True, packed_modules_mapping={} + ) + layer = _make_linear(quant_config, n, k) w = torch.randn((n, k), device="cuda", dtype=torch.bfloat16) / 10 scale = (w.float().abs().max() / FP8_MAX).clamp(min=1e-12) w_fp8 = (w.float() / scale).to(torch.float8_e4m3fn) - layer.weight.data.copy_(w_fp8) - layer.weight_scale.data.fill_(scale) - layer.input_scale.data.fill_(1.0 / FP8_MAX) + # 0-dim scales exercise weight_loader_v2's scalar reshape branch. + load_linear_weights( + layer, + weight=w_fp8, + weight_scale=scale, + input_scale=torch.tensor(1.0 / FP8_MAX, device="cuda"), + ) w_dequant = w_fp8.float() * scale - return method, layer, w_dequant + return layer, w_dequant def test_auto(self): self._check_backend("auto", ["auto"], PER_TENSOR_SHAPES, self._build_layer) diff --git a/test/registered/unit/layers/quantization/test_nvfp4_linear_backends.py b/test/registered/unit/layers/quantization/test_nvfp4_linear_backends.py index 1ca187407..dd11a979a 100644 --- a/test/registered/unit/layers/quantization/test_nvfp4_linear_backends.py +++ b/test/registered/unit/layers/quantization/test_nvfp4_linear_backends.py @@ -1,10 +1,8 @@ """Numerics for the NVFP4 dense-linear GEMM backends (--fp4-gemm-backend). -Runs ModelOptFp4LinearMethod end to end (create_weights -> -process_weights_after_loading -> apply) for each SM100 backend choice and -checks the output against a dequantized-reference matmul. This covers both -the per-backend weight preparation (padding / interleave / TRTLLM shuffle) -and the GEMM kernel dispatch. +Real layer path (ColumnParallelLinear -> weight_loader -> weight processing +-> forward) per SM100 backend vs a dequantized-reference matmul; a merged +two-shard case guards the per-partition scale gathering of fused layers. """ import unittest @@ -15,22 +13,25 @@ from flashinfer import fp4_quantize from sglang.srt.layers.quantization import fp4_utils from sglang.srt.layers.quantization.fp4_utils import Fp4GemmRunnerBackend -from sglang.srt.layers.quantization.modelopt_quant import ( - ModelOptFp4Config, - ModelOptFp4LinearMethod, -) +from sglang.srt.layers.quantization.modelopt_quant import ModelOptFp4Config from sglang.srt.utils import get_device_sm from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.layer_ut_utils import ( + assert_output_close, + init_single_process_dist, + load_linear_weights, + make_tp1_column_parallel_linear, +) +from sglang.test.quant_ref_utils import ( + FLOAT4_E2M1_MAX, + FLOAT8_E4M3_MAX, + dequantize_nvfp4_to_dtype, + quantize_nvfp4_shard, +) from sglang.test.test_utils import CustomTestCase register_cuda_ci(est_time=120, stage="base-b", runner_config="4-gpu-b200") -kE2M1ToFloat = torch.tensor( - [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=torch.float32 -) -FLOAT8_E4M3_MAX = 448.0 -FLOAT4_E2M1_MAX = 6.0 - # (M, N, K). The second shape hits the padding paths: N=160 is not a multiple # of 128 (TRTLLM shuffle pad) and K=336 is neither a multiple of 32 (CUTLASS # K pad) nor K/16 a multiple of 4 (TRTLLM scale pad). @@ -40,96 +41,88 @@ SHAPES = [ (128, 1024, 1024), ] -BACKENDS = [ - "flashinfer_cutedsl", - "flashinfer_cutlass", - "flashinfer_cudnn", - "flashinfer_trtllm", -] +ACT_SCALE = 1.0 / (FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX) -def convert_swizzled_to_linear(a_sf_swizzled: torch.Tensor, m, k, block_size): - m_tiles = (m + 128 - 1) // 128 - f = block_size * 4 - k_tiles = (k + f - 1) // f - tmp = torch.reshape(a_sf_swizzled, (1, m_tiles, k_tiles, 32, 4, 4)) - tmp = torch.permute(tmp, (0, 1, 4, 3, 2, 5)) - out = tmp.reshape(m_tiles * 128, k_tiles * f // block_size) - # Crop the K-tile padding too: k // block_size scale columns, not k. - return out[0:m, 0 : k // block_size] - - -def break_fp4_bytes(a, dtype): - assert a.dtype == torch.uint8 - m, n = a.shape - a_flat = a.flatten() - high = (a_flat & 0xF0) >> 4 - low = a_flat & 0x0F - combined = torch.stack((low, high), dim=1).flatten() - signs = (combined & 0x08).to(torch.bool) - abs_vals = (combined & 0x07).to(torch.long) - kE2M1 = kE2M1ToFloat.to(device=a.device) - values = kE2M1[abs_vals] * torch.where(signs, -1.0, 1.0) - return values.reshape(m, n * 2).to(dtype=dtype) - - -def dequantize_nvfp4_to_dtype( - tensor_fp4, tensor_sf, global_scale, dtype, device, block_size=16 -): - assert tensor_fp4.dtype == torch.uint8 - m, packed_k = tensor_fp4.shape - k = packed_k * 2 - tensor_f32 = break_fp4_bytes(tensor_fp4, torch.float32) - tensor_f32 = tensor_f32.reshape(m, k // block_size, block_size) - tensor_sf = tensor_sf.view(torch.float8_e4m3fn) - tensor_sf = convert_swizzled_to_linear(tensor_sf, m, k, block_size) - tensor_sf_dtype = tensor_sf.to(torch.float32) / global_scale - out = (tensor_f32 * tensor_sf_dtype.unsqueeze(-1)).reshape(m, k) - return out.to(dtype=dtype) - - -def _make_quantized_layer(n: int, k: int, device: str = "cuda"): - """Build a linear layer holding NVFP4 checkpoint-format weights; returns - (method, layer, w_dequant) with w_dequant the fp32 quant->dequant reference.""" +def _make_quantized_layer(n: int, k: int): + """NVFP4 checkpoint-format weights through the real weight_loader.""" quant_config = ModelOptFp4Config( is_checkpoint_nvfp4_serialized=True, group_size=16, use_per_token_activation=False, + packed_modules_mapping={}, ) - method = ModelOptFp4LinearMethod(quant_config) - layer = torch.nn.Module() - method.create_weights( + layer = make_tp1_column_parallel_linear(quant_config, n, k) + + w = torch.randn((n, k), device="cuda", dtype=torch.bfloat16) / 10 + w_q, sf_linear, gs, w_dequant = quantize_nvfp4_shard(w) + load_linear_weights( layer, - input_size_per_partition=k, - output_partition_sizes=[n], + weight=w_q, + weight_scale=sf_linear, + weight_scale_2=(1.0 / gs).clone(), + # Calibrated activation amax stand-in (inputs are randn/10). + input_scale=torch.tensor(ACT_SCALE, device="cuda"), + ) + return layer, w_dequant + + +def _make_merged_layer(n_half: int, k: int): + """Two fused output shards (gate_up_proj) loaded per shard; exercises the + per-partition scale_2 / input_scale gathering that fused-QKV regressions hit.""" + from sglang.srt.layers.linear import MergedColumnParallelLinear + + quant_config = ModelOptFp4Config( + is_checkpoint_nvfp4_serialized=True, + group_size=16, + use_per_token_activation=False, + packed_modules_mapping={"gate_up_proj": ["gate_proj", "up_proj"]}, + ) + layer = MergedColumnParallelLinear( input_size=k, - output_size=n, + output_sizes=[n_half, n_half], + bias=False, params_dtype=torch.bfloat16, - weight_loader=lambda *args, **kwargs: None, - ) - layer = layer.to(device) + quant_config=quant_config, + prefix="model.layers.0.mlp.gate_up_proj", + tp_rank=0, + tp_size=1, + ).cuda() - w = torch.randn((n, k), device=device, dtype=torch.bfloat16) / 10 - w_gs = FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX / w.abs().max().to(torch.float32) - w_q, w_sf_swizzled = fp4_quantize(w, w_gs) - w_sf_linear = convert_swizzled_to_linear( - w_sf_swizzled.view(torch.float8_e4m3fn), n, k, 16 + # process_weights_after_loading collapses shard scale_2 with max() without + # requanting block scales, so shards must share one gs (modelopt fused + # exports ship equal scale_2). + shards = [ + torch.randn((n_half, k), device="cuda", dtype=torch.bfloat16) / 10 + for _ in (0, 1) + ] + shared_gs = ( + FLOAT8_E4M3_MAX + * FLOAT4_E2M1_MAX + / max(w.abs().max().to(torch.float32) for w in shards) ) - w_dequant = dequantize_nvfp4_to_dtype( - w_q, w_sf_swizzled, w_gs, torch.float32, device - ) - - layer.weight.data.copy_(w_q) - layer.weight_scale.data.copy_(w_sf_linear) - layer.weight_scale_2.data.fill_(1.0 / w_gs) - # Calibrated activation amax stand-in (inputs are randn/10). - layer.input_scale.data.fill_(1.0 / (FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX)) - return method, layer, w_dequant + dequants = [] + for shard_id, w in enumerate(shards): + w_q, sf_linear, gs, w_dequant = quantize_nvfp4_shard(w, gs=shared_gs) + load_linear_weights( + layer, + shard_id=shard_id, + weight=w_q, + weight_scale=sf_linear, + weight_scale_2=(1.0 / gs).clone(), + input_scale=torch.tensor(ACT_SCALE, device="cuda"), + ) + dequants.append(w_dequant) + return layer, torch.cat(dequants, dim=0) @unittest.skipIf(get_device_sm() < 100, "NVFP4 dense GEMM backends require SM100+") class TestNvFp4LinearBackends(CustomTestCase): - def _run_backend(self, backend: str): + @classmethod + def setUpClass(cls): + init_single_process_dist() + + def _run_backend(self, backend: str, build_layer=_make_quantized_layer): torch.manual_seed(7) for m, n, k in SHAPES: with self.subTest(backend=backend, shape=(m, n, k)): @@ -138,25 +131,32 @@ class TestNvFp4LinearBackends(CustomTestCase): "FP4_GEMM_RUNNER_BACKEND", Fp4GemmRunnerBackend(backend), ): - method, layer, w_dequant = _make_quantized_layer(n, k) - method.process_weights_after_loading(layer) + layer, w_dequant = build_layer(n, k) + layer.quant_method.process_weights_after_loading(layer) x = torch.randn((m, k), device="cuda", dtype=torch.bfloat16) / 10 - out = method.apply(layer, x) + out, _ = layer(x) + self._assert_matches(layer, x, out, w_dequant) - x_gs = layer.input_scale_inv.data.float() - x_q, x_sf = fp4_quantize(x, x_gs) - x_dequant = dequantize_nvfp4_to_dtype( - x_q, x_sf, x_gs, torch.float32, x.device - ) - ref = x_dequant @ w_dequant.T + def _assert_matches(self, layer, x, out, w_dequant): + x_gs = layer.input_scale_inv.data.float() + x_q, x_sf = fp4_quantize(x, x_gs) + x_dequant = dequantize_nvfp4_to_dtype(x_q, x_sf, x_gs, torch.float32) + ref = x_dequant @ w_dequant.T + assert_output_close(self, out, ref, rtol=5e-2, atol=5e-2) - self.assertEqual(out.shape, (m, n)) - cos = torch.nn.functional.cosine_similarity( - out.float().flatten(), ref.flatten(), dim=0 - ).item() - self.assertGreater(cos, 0.99) - torch.testing.assert_close(out.float(), ref, rtol=5e-2, atol=5e-2) + def test_merged_shards(self): + torch.manual_seed(7) + with mock.patch.object( + fp4_utils, + "FP4_GEMM_RUNNER_BACKEND", + Fp4GemmRunnerBackend("flashinfer_cutedsl"), + ): + layer, w_dequant = _make_merged_layer(256, 512) + layer.quant_method.process_weights_after_loading(layer) + x = torch.randn((16, 512), device="cuda", dtype=torch.bfloat16) / 10 + out, _ = layer(x) + self._assert_matches(layer, x, out, w_dequant) def test_flashinfer_cutedsl(self): self._run_backend("flashinfer_cutedsl") diff --git a/test/registered/unit/layers/quantization/test_nvfp4_moe_backends.py b/test/registered/unit/layers/quantization/test_nvfp4_moe_backends.py index f9834ea34..6b2164b75 100644 --- a/test/registered/unit/layers/quantization/test_nvfp4_moe_backends.py +++ b/test/registered/unit/layers/quantization/test_nvfp4_moe_backends.py @@ -1,13 +1,10 @@ """Numerics for the NVFP4 FusedMoE runner backends (--moe-runner-backend). -Runs the real FusedMoE layer path (construct -> fill NVFP4 checkpoint-format -weights -> process_weights_after_loading -> forward) per backend against a -dequantized torch MoE reference, covering the per-backend weight preparation -(TRTLLM shuffle / CUTLASS swizzle / CuteDSL v2 interleave + MMA blockscales) -and the MoE runner dispatch. Single GPU, tp=ep=1. +Real FusedMoE path (NVFP4 checkpoint shards through the real weight_loader +-> weight processing -> forward) per backend vs a dequantized torch MoE +reference. Single GPU, tp=ep=1. """ -import os import unittest import torch @@ -17,83 +14,30 @@ from sglang.srt.layers.quantization.modelopt_quant import ModelOptFp4Config from sglang.srt.runtime_context import get_context, get_flags, get_parallel from sglang.srt.utils import get_device_sm from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.layer_ut_utils import assert_output_close, init_single_process_dist +from sglang.test.quant_ref_utils import ( + FLOAT4_E2M1_MAX, + FLOAT8_E4M3_MAX, + dequantize_nvfp4_to_dtype, + quantize_nvfp4_shard, +) from sglang.test.test_utils import CustomTestCase register_cuda_ci(est_time=120, stage="base-b", runner_config="4-gpu-b200") E, H, I, TOPK, M = 8, 1024, 1024, 2, 32 -FLOAT8_E4M3_MAX = 448.0 -FLOAT4_E2M1_MAX = 6.0 - -kE2M1ToFloat = torch.tensor( - [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=torch.float32 -) - - -def _init_single_process_dist(): - os.environ.setdefault("MASTER_ADDR", "127.0.0.1") - os.environ.setdefault("MASTER_PORT", "29631") - os.environ.setdefault("RANK", "0") - os.environ.setdefault("WORLD_SIZE", "1") - os.environ.setdefault("LOCAL_RANK", "0") - from sglang.srt.distributed.parallel_state import ( - init_distributed_environment, - initialize_model_parallel, - model_parallel_is_initialized, - ) - - if not torch.distributed.is_initialized(): - init_distributed_environment(world_size=1, rank=0, local_rank=0, backend="gloo") - if not model_parallel_is_initialized(): - initialize_model_parallel( - tensor_model_parallel_size=1, - expert_model_parallel_size=1, - pipeline_model_parallel_size=1, - backend="gloo", - ) - - -def convert_swizzled_to_linear(a_sf_swizzled, m, k, block_size=16): - m_tiles = (m + 128 - 1) // 128 - f = block_size * 4 - k_tiles = (k + f - 1) // f - tmp = torch.reshape(a_sf_swizzled, (1, m_tiles, k_tiles, 32, 4, 4)) - tmp = torch.permute(tmp, (0, 1, 4, 3, 2, 5)) - out = tmp.reshape(m_tiles * 128, k_tiles * f // block_size) - return out[0:m, 0 : k // block_size] - - -def break_fp4_bytes(a): - m, n = a.shape - a_flat = a.flatten() - high = (a_flat & 0xF0) >> 4 - low = a_flat & 0x0F - combined = torch.stack((low, high), dim=1).flatten() - signs = (combined & 0x08).to(torch.bool) - abs_vals = (combined & 0x07).to(torch.long) - kE2M1 = kE2M1ToFloat.to(device=a.device) - values = kE2M1[abs_vals] * torch.where(signs, -1.0, 1.0) - return values.reshape(m, n * 2).to(dtype=torch.float32) - - -def dequant_nvfp4(w_q, sf_swizzled, gs, n, k): - w_f32 = break_fp4_bytes(w_q).reshape(n, k // 16, 16) - sf = convert_swizzled_to_linear(sf_swizzled.view(torch.float8_e4m3fn), n, k) - return (w_f32 * (sf.float() / gs).unsqueeze(-1)).reshape(n, k) @unittest.skipIf(get_device_sm() < 100, "NVFP4 MoE backends require SM100+") class TestNvFp4MoeBackends(CustomTestCase): @classmethod def setUpClass(cls): - _init_single_process_dist() + init_single_process_dist(master_port=29631) torch.set_default_device("cuda") def _run_backend(self, backend: str): - from flashinfer import fp4_quantize - from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE - from sglang.srt.layers.moe.topk import StandardTopKOutput + from sglang.srt.layers.moe.topk import TopKConfig, select_experts torch.manual_seed(7) quant_config = ModelOptFp4Config( @@ -122,93 +66,78 @@ class TestNvFp4MoeBackends(CustomTestCase): gate_up_interleaved=False, ).cuda() - w13_ref = torch.zeros(E, 2 * I, H, dtype=torch.float32, device="cuda") - w2_ref = torch.zeros(E, H, I, dtype=torch.float32, device="cuda") - for e in range(E): - w13 = torch.randn(2 * I, H, dtype=torch.bfloat16, device="cuda") / 10 - w2 = torch.randn(H, I, dtype=torch.bfloat16, device="cuda") / 10 - w13_gs = FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX / w13.abs().max().float() - w2_gs = FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX / w2.abs().max().float() - w13_q, w13_sf = fp4_quantize(w13, w13_gs) - w2_q, w2_sf = fp4_quantize(w2, w2_gs) - layer.w13_weight.data[e].copy_(w13_q) - layer.w2_weight.data[e].copy_(w2_q) - layer.w13_weight_scale.data[e].copy_( - convert_swizzled_to_linear( - w13_sf.view(torch.float8_e4m3fn), 2 * I, H - ) - ) - layer.w2_weight_scale.data[e].copy_( - convert_swizzled_to_linear(w2_sf.view(torch.float8_e4m3fn), H, I) - ) - layer.w13_weight_scale_2.data[e].fill_(1.0 / w13_gs) - layer.w2_weight_scale_2.data[e].fill_(1.0 / w2_gs) - w13_ref[e] = dequant_nvfp4(w13_q, w13_sf, w13_gs, 2 * I, H) - w2_ref[e] = dequant_nvfp4(w2_q, w2_sf, w2_gs, H, I) + # Checkpoint-format shards through the real weight_loader; + # gate/up placement stays the loader's job. + refs = { + "w1": torch.zeros(E, I, H, dtype=torch.float32, device="cuda"), + "w3": torch.zeros(E, I, H, dtype=torch.float32, device="cuda"), + "w2": torch.zeros(E, H, I, dtype=torch.float32, device="cuda"), + } act_scale = 1.0 / (FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX) - layer.w13_input_scale.data.fill_(act_scale) - layer.w2_input_scale.data.fill_(act_scale) - + for e in range(E): + for shard_id in ("w1", "w3", "w2"): + rows, cols = (I, H) if shard_id in ("w1", "w3") else (H, I) + w = ( + torch.randn(rows, cols, dtype=torch.bfloat16, device="cuda") + / 10 + ) + w_q, sf_linear, gs, w_dequant = quantize_nvfp4_shard(w) + prefix = "w13" if shard_id in ("w1", "w3") else "w2" + for suffix, loaded in ( + ("weight", w_q), + ("weight_scale", sf_linear), + ("weight_scale_2", (1.0 / gs).clone()), + ("input_scale", torch.tensor(act_scale, device="cuda")), + ): + name = f"{prefix}_{suffix}" + param = getattr(layer, name) + layer.weight_loader( + param, loaded, name, shard_id=shard_id, expert_id=e + ) + refs[shard_id][e] = w_dequant layer.quant_method.process_weights_after_loading(layer) x = torch.randn(M, H, dtype=torch.bfloat16, device="cuda") / 10 router_logits = torch.randn(M, E, dtype=torch.float32, device="cuda") - weights = torch.softmax(router_logits, dim=-1) - topk_weights, topk_ids = torch.topk(weights, TOPK, dim=-1) - topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True) - - out = layer.forward( - x, - StandardTopKOutput( - topk_weights=topk_weights, - topk_ids=topk_ids.to(torch.int32), - router_logits=router_logits, - ), + # Route through the real topk compute path, not a hand-rolled one. + topk_output = select_experts( + hidden_states=x, + router_logits=router_logits, + topk_config=TopKConfig(top_k=TOPK, renormalize=True), ) + topk_weights, topk_ids = topk_output.topk_weights, topk_output.topk_ids + + out = layer.forward(x, topk_output) if not isinstance(out, torch.Tensor): out = out[0] if isinstance(out, tuple) else out.hidden_states ref = self._torch_moe_reference( - layer, x, topk_weights, topk_ids, w13_ref, w2_ref + x, topk_weights, topk_ids, refs["w1"], refs["w3"], refs["w2"] ) - - self.assertEqual(out.shape, (M, H)) - cos = torch.nn.functional.cosine_similarity( - out.float().flatten(), ref.flatten(), dim=0 - ).item() - self.assertGreater(cos, 0.99) + assert_output_close(self, out, ref) @staticmethod - def _torch_moe_reference(layer, x, topk_weights, topk_ids, w13_ref, w2_ref): + def _torch_moe_reference(x, topk_weights, topk_ids, w1_ref, w3_ref, w2_ref): from flashinfer import fp4_quantize def quant_roundtrip(t2d, gs): q, sf = fp4_quantize(t2d.to(torch.bfloat16), gs) - return dequant_nvfp4(q, sf, gs, t2d.shape[0], t2d.shape[1]) + return dequantize_nvfp4_to_dtype(q, sf, gs, torch.float32) - # The kernels quantize the input and the GEMM1->GEMM2 intermediate to - # NVFP4; mirror both round trips or the comparison carries ~7% noise. + # Mirror the kernels' two fp4 activation round trips (input and + # GEMM1->GEMM2) or the comparison carries ~7% noise. The reference + # stays in checkpoint semantics (w1=gate, w3=up). act_gs = torch.tensor( FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX, dtype=torch.float32, device="cuda" ) - # TRTLLM consumes w13 as [up; gate] (GEMM1 scales are applied on that - # assumption); CUTLASS / CuteDSL-v2 load up first as well via - # load_up_proj_weight_first. - up_first = ( - layer.quant_method.load_up_proj_weight_first - or layer.quant_method.enable_flashinfer_trtllm_moe - ) x_dq = quant_roundtrip(x.float(), act_gs) m, h = x.shape ref = torch.zeros(m, h, dtype=torch.float32, device="cuda") for t in range(m): for j in range(topk_ids.shape[1]): e = int(topk_ids[t, j]) - gu = x_dq[t] @ w13_ref[e].T - if up_first: - up, gate = gu[:I], gu[I:] - else: - gate, up = gu[:I], gu[I:] + gate = x_dq[t] @ w1_ref[e].T + up = x_dq[t] @ w3_ref[e].T act = torch.nn.functional.silu(gate) * up act_dq = quant_roundtrip(act.unsqueeze(0), act_gs)[0] ref[t] += float(topk_weights[t, j]) * (act_dq @ w2_ref[e].T) diff --git a/test/registered/unit/models/test_zaya_cca.py b/test/registered/unit/models/test_zaya_cca.py index d3740cea2..e9cfadd71 100644 --- a/test/registered/unit/models/test_zaya_cca.py +++ b/test/registered/unit/models/test_zaya_cca.py @@ -21,7 +21,6 @@ GPU dependency. State is stored in a mock centralized pool that mirrors the ``HybridReqToTokenPool`` / ``MambaPool`` interface used at serving time. """ -import os import unittest from contextlib import contextmanager from dataclasses import dataclass @@ -31,55 +30,16 @@ from typing import List, Optional import torch from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.layer_ut_utils import init_single_process_dist from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=30, suite="base-a-test-cpu") def _ensure_dist_initialized() -> None: - """Set up a minimal single-rank gloo distributed environment plus the - SGLang model-parallel groups (TP=1, PP=1, EP=1). The CCA module reads - ``get_tensor_model_parallel_rank()`` / ``get_tensor_model_parallel_world_size()`` - inside ``__init__`` to size its head-parallel projections, so the world - group and model parallel groups must both be initialized before any CCA - construction. - """ - os.environ.setdefault("MASTER_ADDR", "127.0.0.1") - os.environ.setdefault("MASTER_PORT", "29632") - os.environ.setdefault("RANK", "0") - os.environ.setdefault("WORLD_SIZE", "1") - os.environ.setdefault("LOCAL_RANK", "0") - - from sglang.srt.distributed.parallel_state import ( - init_distributed_environment, - initialize_model_parallel, - model_parallel_is_initialized, - ) - - if not torch.distributed.is_initialized(): - init_distributed_environment( - world_size=1, - rank=0, - local_rank=0, - backend="gloo", - ) - - if not model_parallel_is_initialized(): - # Pass arguments as kwargs because ``ensure_model_parallel_initialized`` - # forwards positional ``backend`` into the ``attention_data_parallel_size`` - # slot of ``initialize_model_parallel``, which then explodes on - # ``int // str``. Using kwargs avoids that footgun. - initialize_model_parallel( - tensor_model_parallel_size=1, - expert_model_parallel_size=1, - pipeline_model_parallel_size=1, - backend="gloo", - ) - - -# --------------------------------------------------------------------------- -# Mock centralized pool -# --------------------------------------------------------------------------- + """CCA reads the TP rank / world size inside ``__init__`` to size its + head-parallel projections, so the groups must exist before construction.""" + init_single_process_dist() @dataclass(frozen=True) @@ -202,11 +162,6 @@ def _mock_pool_context(pool: _MockReqToTokenPool): set_forward_context(prev) -# --------------------------------------------------------------------------- -# Helper factories -# --------------------------------------------------------------------------- - - def _make_forward_batch( *, is_decode: bool,