Select DeepGEMM standard layouts by memory budget (#33474)
Co-authored-by: Chunan Zeng <zcnrex@gmail.com>
This commit is contained in:
@@ -733,6 +733,8 @@ class Envs:
|
|||||||
|
|
||||||
# DeepGemm
|
# DeepGemm
|
||||||
SGLANG_ENABLE_JIT_DEEPGEMM = EnvBool(True)
|
SGLANG_ENABLE_JIT_DEEPGEMM = EnvBool(True)
|
||||||
|
SGLANG_DEEPGEMM_STANDARD_LAYOUT = EnvStr("auto")
|
||||||
|
SGLANG_DEEPGEMM_MASKED_MEMORY_BUDGET_FRACTION = EnvFloat(0.25)
|
||||||
# Cap the DeepGEMM masked grouped-GEMM per-expert padded capacity at
|
# Cap the DeepGEMM masked grouped-GEMM per-expert padded capacity at
|
||||||
# round_up(max(masked_m), 256) instead of round_up(rank_tokens, 256):
|
# round_up(max(masked_m), 256) instead of round_up(rank_tokens, 256):
|
||||||
# shrinks the [num_local_experts, m, *] MoE intermediates ~4x under
|
# shrinks the [num_local_experts, m, *] MoE intermediates ~4x under
|
||||||
|
|||||||
@@ -73,6 +73,7 @@ else:
|
|||||||
|
|
||||||
|
|
||||||
_DEEPGEMM_ON_H20 = get_bool_env_var("SGLANG_DEEPGEMM_ON_H20")
|
_DEEPGEMM_ON_H20 = get_bool_env_var("SGLANG_DEEPGEMM_ON_H20")
|
||||||
|
_masked_standard_layout_memory_budget_bytes: Optional[int] = None
|
||||||
|
|
||||||
|
|
||||||
# TODO(kaixih@nvidia): ideally we should merge this logic into
|
# TODO(kaixih@nvidia): ideally we should merge this logic into
|
||||||
@@ -100,11 +101,91 @@ def copy_list_to_gpu_no_ce(arr: List[int]):
|
|||||||
return tensor_gpu
|
return tensor_gpu
|
||||||
|
|
||||||
|
|
||||||
def _should_use_masked_standard_layout(runner_config: MoeRunnerConfig) -> bool:
|
def set_masked_standard_layout_memory_budget(
|
||||||
"""Use masked GEMM when expert parallelism keeps its buffer small."""
|
available_memory_bytes: int,
|
||||||
|
) -> int:
|
||||||
|
"""Cache the masked-layout share of free non-static device memory."""
|
||||||
|
global _masked_standard_layout_memory_budget_bytes
|
||||||
|
fraction = envs.SGLANG_DEEPGEMM_MASKED_MEMORY_BUDGET_FRACTION.get()
|
||||||
|
if not 0.0 < fraction <= 1.0:
|
||||||
|
raise ValueError(
|
||||||
|
"SGLANG_DEEPGEMM_MASKED_MEMORY_BUDGET_FRACTION must be in (0, 1]"
|
||||||
|
)
|
||||||
|
_masked_standard_layout_memory_budget_bytes = int(available_memory_bytes * fraction)
|
||||||
|
return _masked_standard_layout_memory_budget_bytes
|
||||||
|
|
||||||
|
|
||||||
|
def _estimate_masked_standard_layout_peak_bytes(
|
||||||
|
runner_config: MoeRunnerConfig,
|
||||||
|
quant_info: DeepGemmMoeQuantInfo,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
) -> int:
|
||||||
|
padded_m = (hidden_states.shape[0] // 256 + 1) * 256
|
||||||
|
activation_dtype = (
|
||||||
|
torch.bfloat16
|
||||||
|
if quant_info.w13_weight.dtype == torch.bfloat16
|
||||||
|
else torch.float8_e4m3fn
|
||||||
|
)
|
||||||
|
hidden_size = hidden_states.shape[1]
|
||||||
|
gateup_size = quant_info.w13_weight.shape[1]
|
||||||
|
gateup_row_bytes = gateup_size * torch.bfloat16.itemsize
|
||||||
|
down_output_row_bytes = quant_info.w2_weight.shape[1] * torch.bfloat16.itemsize
|
||||||
|
input_row_bytes = hidden_size * activation_dtype.itemsize
|
||||||
|
down_input_row_bytes = gateup_size // 2 * activation_dtype.itemsize
|
||||||
|
|
||||||
|
if activation_dtype == torch.bfloat16:
|
||||||
|
input_scale_row_bytes = 0
|
||||||
|
down_scale_row_bytes = 0
|
||||||
|
else:
|
||||||
|
block_k = quant_info.block_shape[1] if quant_info.block_shape else 128
|
||||||
|
packed_scales = quant_info.use_mxfp8 or deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0
|
||||||
|
scale_item_bytes = (
|
||||||
|
torch.uint8.itemsize if packed_scales else torch.float32.itemsize
|
||||||
|
)
|
||||||
|
input_scale_row_bytes = ceil_div(hidden_size, block_k) * scale_item_bytes
|
||||||
|
down_scale_row_bytes = ceil_div(gateup_size // 2, block_k) * scale_item_bytes
|
||||||
|
|
||||||
|
peak_row_bytes = max(
|
||||||
|
input_row_bytes + input_scale_row_bytes + gateup_row_bytes,
|
||||||
|
gateup_row_bytes + down_input_row_bytes + down_scale_row_bytes,
|
||||||
|
down_input_row_bytes + down_scale_row_bytes + down_output_row_bytes,
|
||||||
|
)
|
||||||
|
return runner_config.num_local_experts * padded_m * peak_row_bytes
|
||||||
|
|
||||||
|
|
||||||
|
def _should_use_masked_standard_layout(
|
||||||
|
runner_config: MoeRunnerConfig,
|
||||||
|
quant_info: DeepGemmMoeQuantInfo,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
) -> bool:
|
||||||
|
mode = envs.SGLANG_DEEPGEMM_STANDARD_LAYOUT.get().lower()
|
||||||
|
if mode not in ("auto", "masked", "compact"):
|
||||||
|
raise ValueError(
|
||||||
|
"SGLANG_DEEPGEMM_STANDARD_LAYOUT must be one of: auto, masked, compact"
|
||||||
|
)
|
||||||
|
if mode != "auto":
|
||||||
|
return mode == "masked"
|
||||||
|
|
||||||
|
global _masked_standard_layout_memory_budget_bytes
|
||||||
|
if _masked_standard_layout_memory_budget_bytes is None:
|
||||||
|
# Serving sets an all-rank budget before capture. Direct eager callers
|
||||||
|
# fall back to this rank's free memory without querying inside capture.
|
||||||
|
# Import lazily to avoid a module-initialization cycle through
|
||||||
|
# runner_utils -> DeepEP -> MoE -> this module.
|
||||||
|
from sglang.srt.model_executor.runner_utils.capture_mode import (
|
||||||
|
get_is_capture_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
if get_is_capture_mode():
|
||||||
|
return False
|
||||||
|
free_memory, _ = torch.cuda.mem_get_info(hidden_states.device)
|
||||||
|
set_masked_standard_layout_memory_budget(free_memory)
|
||||||
|
|
||||||
return (
|
return (
|
||||||
runner_config.num_experts > runner_config.num_local_experts
|
_estimate_masked_standard_layout_peak_bytes(
|
||||||
and runner_config.num_local_experts <= 32
|
runner_config, quant_info, hidden_states
|
||||||
|
)
|
||||||
|
<= _masked_standard_layout_memory_budget_bytes
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -775,7 +856,7 @@ def pre_permute_standard_to_deep_gemm(
|
|||||||
|
|
||||||
topk_weights, topk_ids = topk_weights, topk_ids
|
topk_weights, topk_ids = topk_weights, topk_ids
|
||||||
|
|
||||||
if _should_use_masked_standard_layout(runner_config):
|
if _should_use_masked_standard_layout(runner_config, quant_info, hidden_states):
|
||||||
output_dtype = (
|
output_dtype = (
|
||||||
torch.bfloat16
|
torch.bfloat16
|
||||||
if quant_info.w13_weight.dtype == torch.bfloat16
|
if quant_info.w13_weight.dtype == torch.bfloat16
|
||||||
|
|||||||
@@ -1073,20 +1073,28 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
), "cutlass_fp8 MoE requires SM90, SM100, or SM120 GPUs"
|
), "cutlass_fp8 MoE requires SM90, SM100, or SM120 GPUs"
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def is_deepgemm_moe_runner_backend_enabled() -> bool:
|
def is_deepgemm_moe_runner_backend_enabled(
|
||||||
|
moe_runner_backend=None, moe_a2a_backend=None
|
||||||
|
) -> bool:
|
||||||
"""Check if MoE will actually use DeepGEMM runner for FP8."""
|
"""Check if MoE will actually use DeepGEMM runner for FP8."""
|
||||||
from sglang.srt.layers import deep_gemm_wrapper
|
|
||||||
from sglang.srt.layers.moe.utils import get_moe_a2a_backend
|
from sglang.srt.layers.moe.utils import get_moe_a2a_backend
|
||||||
|
|
||||||
|
if moe_runner_backend is None:
|
||||||
moe_runner_backend = get_moe_runner_backend()
|
moe_runner_backend = get_moe_runner_backend()
|
||||||
if moe_runner_backend.is_deep_gemm():
|
if moe_runner_backend.is_deep_gemm():
|
||||||
return True
|
return True
|
||||||
if moe_runner_backend.is_auto():
|
if moe_runner_backend.is_auto():
|
||||||
return deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM and (
|
if moe_a2a_backend is None:
|
||||||
get_moe_a2a_backend().is_deepep()
|
moe_a2a_backend = get_moe_a2a_backend()
|
||||||
or get_moe_a2a_backend().is_mooncake()
|
if not (
|
||||||
or get_moe_a2a_backend().is_nixl()
|
moe_a2a_backend.is_deepep()
|
||||||
)
|
or moe_a2a_backend.is_mooncake()
|
||||||
|
or moe_a2a_backend.is_nixl()
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
from sglang.srt.layers import deep_gemm_wrapper
|
||||||
|
|
||||||
|
return deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -8,9 +8,11 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
import msgspec
|
import msgspec
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import ModelImpl
|
from sglang.srt.configs.model_config import ModelImpl
|
||||||
|
from sglang.srt.distributed import get_world_group
|
||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
prealloc_symmetric_memory_pool,
|
prealloc_symmetric_memory_pool,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner
|
from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner
|
||||||
from sglang.srt.hardware_backend.xpu.graph_runner.xpu_graph_runner import XPUGraphRunner
|
from sglang.srt.hardware_backend.xpu.graph_runner.xpu_graph_runner import XPUGraphRunner
|
||||||
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
||||||
@@ -125,6 +127,57 @@ def capture_cuda_graphs(
|
|||||||
# batch.
|
# batch.
|
||||||
eager_runner = EagerRunner(model_runner)
|
eager_runner = EagerRunner(model_runner)
|
||||||
|
|
||||||
|
if model_runner.is_draft_worker:
|
||||||
|
moe_runner_backend = (
|
||||||
|
model_runner.server_args.speculative_moe_runner_backend
|
||||||
|
or model_runner.server_args.moe_runner_backend
|
||||||
|
)
|
||||||
|
moe_a2a_backend = (
|
||||||
|
model_runner.server_args.speculative_moe_a2a_backend
|
||||||
|
or model_runner.server_args.moe_a2a_backend
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
moe_runner_backend = model_runner.server_args.moe_runner_backend
|
||||||
|
moe_a2a_backend = model_runner.server_args.moe_a2a_backend
|
||||||
|
|
||||||
|
uses_deep_gemm_moe_runner = moe_runner_backend == "deep_gemm"
|
||||||
|
if moe_runner_backend == "auto" and model_runner.model_config.quantization in (
|
||||||
|
"fp8",
|
||||||
|
"mxfp8",
|
||||||
|
):
|
||||||
|
from sglang.srt.layers.moe.utils import MoeA2ABackend, MoeRunnerBackend
|
||||||
|
from sglang.srt.layers.quantization.fp8 import Fp8MoEMethod
|
||||||
|
|
||||||
|
uses_deep_gemm_moe_runner = Fp8MoEMethod.is_deepgemm_moe_runner_backend_enabled(
|
||||||
|
MoeRunnerBackend(moe_runner_backend),
|
||||||
|
MoeA2ABackend(moe_a2a_backend),
|
||||||
|
)
|
||||||
|
|
||||||
|
if (
|
||||||
|
model_runner.device == "cuda"
|
||||||
|
and envs.SGLANG_DEEPGEMM_STANDARD_LAYOUT.get().lower() == "auto"
|
||||||
|
and uses_deep_gemm_moe_runner
|
||||||
|
):
|
||||||
|
from sglang.srt.layers.moe.moe_runner.deep_gemm import (
|
||||||
|
set_masked_standard_layout_memory_budget,
|
||||||
|
)
|
||||||
|
|
||||||
|
world_group = get_world_group()
|
||||||
|
available_memory_gb = get_available_gpu_memory(
|
||||||
|
model_runner.device,
|
||||||
|
model_runner.gpu_id,
|
||||||
|
distributed=world_group.world_size > 1,
|
||||||
|
cpu_group=world_group.cpu_group,
|
||||||
|
)
|
||||||
|
budget_bytes = set_masked_standard_layout_memory_budget(
|
||||||
|
int(available_memory_gb * (1 << 30))
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"DeepGEMM masked layout budget: %.2f GiB from %.2f GiB free.",
|
||||||
|
budget_bytes / (1 << 30),
|
||||||
|
available_memory_gb,
|
||||||
|
)
|
||||||
|
|
||||||
# cuda-graph capture: prefill before decode, so both coalesce onto the
|
# cuda-graph capture: prefill before decode, so both coalesce onto the
|
||||||
# eager buffer allocated above. (capture_prefill_graph routes prefill
|
# eager buffer allocated above. (capture_prefill_graph routes prefill
|
||||||
# to the eager runner when the prefill graph is disabled.)
|
# to the eager runner when the prefill graph is disabled.)
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from sglang.kernels.ops.quantization.minimax_quant_ue8m0 import (
|
|||||||
per_token_quant_fp8_ue8m0,
|
per_token_quant_fp8_ue8m0,
|
||||||
per_token_quant_fp8_ue8m0_scatter,
|
per_token_quant_fp8_ue8m0_scatter,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
|
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
|
||||||
from sglang.srt.layers.moe.moe_runner.deep_gemm import (
|
from sglang.srt.layers.moe.moe_runner.deep_gemm import (
|
||||||
DeepGemmMoeQuantInfo,
|
DeepGemmMoeQuantInfo,
|
||||||
@@ -158,6 +159,37 @@ def test_compact_all_tokens_uses_tight_routing_independent_bound(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_standard_layout_auto_memory_policy(monkeypatch):
|
||||||
|
config = MoeRunnerConfig(
|
||||||
|
num_experts=512,
|
||||||
|
num_local_experts=512,
|
||||||
|
hidden_size=4096,
|
||||||
|
intermediate_size_per_partition=256,
|
||||||
|
top_k=8,
|
||||||
|
)
|
||||||
|
quant_info = DeepGemmMoeQuantInfo(
|
||||||
|
w13_weight=torch.empty((1, 512, 1), dtype=torch.float8_e4m3fn),
|
||||||
|
w2_weight=torch.empty((1, 4096, 1), dtype=torch.float8_e4m3fn),
|
||||||
|
use_fp8=True,
|
||||||
|
block_shape=[128, 128],
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
deep_gemm_runner,
|
||||||
|
"_masked_standard_layout_memory_budget_bytes",
|
||||||
|
int(42.5 * (1 << 30)),
|
||||||
|
)
|
||||||
|
|
||||||
|
with envs.SGLANG_DEEPGEMM_STANDARD_LAYOUT.override("auto"):
|
||||||
|
for num_tokens, expected in ((8192, True), (16384, False)):
|
||||||
|
hidden_states = torch.empty((num_tokens, 4096), device="meta")
|
||||||
|
assert (
|
||||||
|
deep_gemm_runner._should_use_masked_standard_layout(
|
||||||
|
config, quant_info, hidden_states
|
||||||
|
)
|
||||||
|
is expected
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("weight_dtype", ["fp8", "bf16"])
|
@pytest.mark.parametrize("weight_dtype", ["fp8", "bf16"])
|
||||||
def test_standard_masked_runner_matches_compact_end_to_end(monkeypatch, weight_dtype):
|
def test_standard_masked_runner_matches_compact_end_to_end(monkeypatch, weight_dtype):
|
||||||
"""Exercise both production grouped GEMMs through the standard path."""
|
"""Exercise both production grouped GEMMs through the standard path."""
|
||||||
@@ -252,9 +284,9 @@ def test_standard_masked_runner_matches_compact_end_to_end(monkeypatch, weight_d
|
|||||||
topk_output=(topk_weights, topk_ids, None),
|
topk_output=(topk_weights, topk_ids, None),
|
||||||
)
|
)
|
||||||
|
|
||||||
def run_with_num_experts(num_experts):
|
def run_with_layout(layout):
|
||||||
config = MoeRunnerConfig(
|
config = MoeRunnerConfig(
|
||||||
num_experts=num_experts,
|
num_experts=8,
|
||||||
num_local_experts=num_local_experts,
|
num_local_experts=num_local_experts,
|
||||||
hidden_size=hidden,
|
hidden_size=hidden,
|
||||||
intermediate_size_per_partition=intermediate,
|
intermediate_size_per_partition=intermediate,
|
||||||
@@ -264,6 +296,7 @@ def test_standard_masked_runner_matches_compact_end_to_end(monkeypatch, weight_d
|
|||||||
inplace=False,
|
inplace=False,
|
||||||
)
|
)
|
||||||
running_state = {}
|
running_state = {}
|
||||||
|
with envs.SGLANG_DEEPGEMM_STANDARD_LAYOUT.override(layout):
|
||||||
runner_input = pre_permute_standard_to_deep_gemm(
|
runner_input = pre_permute_standard_to_deep_gemm(
|
||||||
dispatch_output,
|
dispatch_output,
|
||||||
quant_info,
|
quant_info,
|
||||||
@@ -288,10 +321,10 @@ def test_standard_masked_runner_matches_compact_end_to_end(monkeypatch, weight_d
|
|||||||
)
|
)
|
||||||
|
|
||||||
compact_is_masked, compact_all_tokens, compact_m_indices, compact_output = (
|
compact_is_masked, compact_all_tokens, compact_m_indices, compact_output = (
|
||||||
run_with_num_experts(num_local_experts)
|
run_with_layout("compact")
|
||||||
)
|
)
|
||||||
masked_is_masked, masked_all_tokens, masked_m_indices, masked_output = (
|
masked_is_masked, masked_all_tokens, masked_m_indices, masked_output = (
|
||||||
run_with_num_experts(8)
|
run_with_layout("masked")
|
||||||
)
|
)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user