Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
217 changes: 208 additions & 9 deletions tests/unit_tests/extension/test_sglang_extension.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,19 +3,26 @@
from __future__ import annotations

import dataclasses
import os
import sys
import types
from contextlib import contextmanager
from unittest.mock import patch

import pytest
import torch

import megatron.core.parallel_state as parallel_state
from miles_megatron_plugins.true_on_policy.sglang_backend import (
QWEN3_DENSE_SGLANG_MATH,
QWEN3_DENSE_TRUE_ON_POLICY_V1,
QWEN3_MOE_SGLANG_MATH,
QWEN3_MOE_TRUE_ON_POLICY_V1,
MegatronTrueOnPolicyRuntimePolicy,
SGLangColumnParallelLinear,
SGLangCoreAttention,
SGLangFinalRMSNorm,
SGLangGroupedMLP,
SGLangNorm,
SGLangQKRMSNorm,
SGLangRowParallelLinear,
Expand All @@ -27,10 +34,16 @@
resolve_true_on_policy_runtime_policy,
)
from miles_megatron_plugins.true_on_policy.contracts import get_true_on_policy_contract
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_decoder_layer_specs
from megatron.core.models.gpt.gpt_layer_specs import (
get_gpt_decoder_layer_specs,
get_gpt_layer_local_spec,
)
from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear
from megatron.core.tensor_parallel.layers import linear_with_grad_accumulation_and_async_allreduce
from miles_megatron_plugins.true_on_policy.matmul import _sglang_row_parallel_matmul, sglang_reference_matmul
from miles_megatron_plugins.true_on_policy.matmul import (
_sglang_row_parallel_matmul,
sglang_reference_matmul,
)
from megatron.core.models.common.embeddings.rope_utils import apply_rotary_pos_emb
from megatron.core.transformer.custom_layers.batch_invariant_kernels import (
matmul_persistent,
Expand Down Expand Up @@ -183,6 +196,8 @@ def test_contract_object_owns_megatron_runtime_policy_values():
assert contract.schema.model_family == "qwen3_dense"
assert policy.contract_name == QWEN3_DENSE_TRUE_ON_POLICY_V1
assert policy.enabled
assert policy.requires_kernel(QWEN3_DENSE_SGLANG_MATH)
assert not policy.requires_kernel(QWEN3_MOE_SGLANG_MATH)
assert policy.use_sglang_backend
assert policy.batch_invariant_mode
assert policy.disable_rope_fusion
Expand All @@ -199,6 +214,72 @@ def test_contract_object_owns_megatron_runtime_policy_values():
assert not policy.use_ulysses_cp_recompute_fallback


def test_qwen3_moe_contract_splits_tp_cp_and_ep_policy_values():
contract = get_true_on_policy_contract(QWEN3_MOE_TRUE_ON_POLICY_V1)
config = _make_config(
batch_invariant_mode=True,
attention_backend=AttnBackend.flash,
context_parallel_size=2,
cp_comm_type="a2a",
expert_model_parallel_size=4,
num_moe_experts=128,
true_on_policy_contract=QWEN3_MOE_TRUE_ON_POLICY_V1,
)

policy = contract.policy_for(config)

assert contract.schema.model_family == "qwen3_moe"
assert policy.contract_name == QWEN3_MOE_TRUE_ON_POLICY_V1
assert policy.requires_kernel(QWEN3_DENSE_SGLANG_MATH)
assert policy.requires_kernel(QWEN3_MOE_SGLANG_MATH)
assert policy.use_sglang_backend
assert policy.cp_layout == "ulysses_a2a"
assert policy.deterministic_moe_routing
assert policy.moe_topk_tiebreak == "stable_sort"
assert policy.ep_invariant_moe
assert policy.deterministic_moe_dispatch


def test_true_on_policy_moe_local_spec_keeps_pre_mlp_layernorm_checkpoint_key():
dense_spec = get_gpt_layer_local_spec(normalization="RMSNorm", use_true_on_policy_backend=True)
moe_spec = get_gpt_layer_local_spec(
num_experts=128,
moe_grouped_gemm=True,
normalization="RMSNorm",
use_true_on_policy_backend=True,
)

dense_key_map = dense_spec.submodules.sharded_state_dict_keys_map
moe_key_map = moe_spec.submodules.sharded_state_dict_keys_map

assert dense_key_map["pre_mlp_layernorm."] == "mlp.linear_fc1.layer_norm_"
assert "pre_mlp_layernorm." not in moe_key_map
assert "mlp.linear_fc1.layer_norm_" not in moe_key_map.values()


def test_sglang_moe_ep_tree_all_reduce_uses_sglang_fixed_tree(monkeypatch):
from miles_megatron_plugins.true_on_policy.moe import sglang_moe_ep_tree_all_reduce

calls = []
tp_invariant_ops = types.ModuleType("sglang.srt.tp_invariant_ops")

def fake_tree_all_reduce_sum(input_, device_group=None):
calls.append((input_, device_group))
return input_ + 1

tp_invariant_ops.tree_all_reduce_sum = fake_tree_all_reduce_sum
monkeypatch.setitem(sys.modules, "sglang", types.ModuleType("sglang"))
monkeypatch.setitem(sys.modules, "sglang.srt", types.ModuleType("sglang.srt"))
monkeypatch.setitem(sys.modules, "sglang.srt.tp_invariant_ops", tp_invariant_ops)

ep_group = object()
input_ = torch.zeros(2, 3)
output = sglang_moe_ep_tree_all_reduce(input_, ep_group)

torch.testing.assert_close(output, torch.ones_like(input_))
assert calls == [(input_, ep_group)]


def test_qwen3_dense_contract_only_marks_ulysses_a2a_as_cp_layout():
contract = get_true_on_policy_contract(QWEN3_DENSE_TRUE_ON_POLICY_V1)

Expand All @@ -221,6 +302,8 @@ def test_missing_true_on_policy_contract_uses_default_policy():

assert policy.contract_name is None
assert not policy.enabled
assert not policy.required_kernel_contracts
assert not policy.requires_kernel(QWEN3_DENSE_SGLANG_MATH)


def test_invalid_true_on_policy_contract_is_rejected_by_config():
Expand Down Expand Up @@ -628,23 +711,36 @@ def test_sglang_qk_rmsnorm_matches_source_truth_dtype_boundary():


def test_sglang_final_rmsnorm_matches_source_truth_dtype_boundary():
config = _make_config(normalization="RMSNorm")
config = _make_config(normalization="RMSNorm", params_dtype=torch.bfloat16)
norm = SGLangFinalRMSNorm(config=config, hidden_size=4, eps=1e-6)
x = torch.randn(2, 3, 4, dtype=torch.bfloat16)
residual = torch.randn(2, 3, 4, dtype=torch.bfloat16)
x = torch.randn(2, 3, 4, dtype=torch.bfloat16).float()
residual = torch.randn(2, 3, 4, dtype=torch.float32)

actual, actual_residual = norm(x, residual)
x_with_residual = x + residual
x_float = x_with_residual.float()
expected = x_float * torch.rsqrt(x_float.pow(2).mean(-1, keepdim=True) + norm.eps)
expected = norm.weight.float() * expected.to(torch.bfloat16)
expected = norm.weight.to(torch.bfloat16) * expected.to(torch.bfloat16)

assert norm.weight.dtype == torch.float32
assert actual.dtype == torch.float32
assert actual.dtype == torch.bfloat16
torch.testing.assert_close(actual, expected)
torch.testing.assert_close(actual_residual, x_with_residual)


def test_sglang_exact_forward_native_grad_wrapper_preserves_exact_value():
from miles_megatron_plugins.true_on_policy.norm import _with_native_grad

exact = torch.tensor([1.0], dtype=torch.float32)
native = torch.tensor([1.0e8], dtype=torch.float32, requires_grad=True)

actual = _with_native_grad(exact, native)

torch.testing.assert_close(actual.detach(), exact, atol=0.0, rtol=0.0)
actual.sum().backward()
torch.testing.assert_close(native.grad, torch.ones_like(native))


def test_sglang_bias_dropout_add_keeps_residual_add_in_float():
x = torch.randn(2, 3, 4, dtype=torch.bfloat16)
residual = torch.randn(2, 3, 4, dtype=torch.float32)
Expand All @@ -663,7 +759,7 @@ def test_sglang_rmsnorm_rejects_zero_centered_gamma():
SGLangNorm(config=config, hidden_size=4, zero_centered_gamma=True)


def test_sglang_spec_provider_grouped_mlp_fallback():
def test_sglang_spec_provider_grouped_mlp_uses_sglang_no_grad_surface():
provider = SGLangSpecProvider()

module, submodules = provider.grouped_mlp_modules(
Expand All @@ -676,10 +772,113 @@ def test_sglang_spec_provider_grouped_mlp_fallback():
grouped_module, grouped_submodules = provider.grouped_mlp_modules(
moe_use_grouped_gemm=True, moe_use_legacy_grouped_gemm=False
)
assert grouped_module.__name__ == "GroupedMLP"
assert grouped_module is SGLangGroupedMLP
assert grouped_submodules is None


def test_sglang_grouped_mlp_weight_views_match_megatron_grouped_layout():
config = _make_config(hidden_size=2, moe_ffn_hidden_size=3, num_moe_experts=2)
layer = object.__new__(SGLangGroupedMLP)
layer.config = config
layer.num_local_experts = 2
layer.weight1 = torch.arange(2 * 2 * 2 * 3, dtype=torch.bfloat16).reshape(2, 12)
layer.weight2 = torch.arange(2 * 3 * 2, dtype=torch.bfloat16).reshape(6, 2)

expected_w13 = layer.weight1.view(2, 2, 6).permute(0, 2, 1).contiguous()
expected_w2 = layer.weight2.view(2, 3, 2).permute(0, 2, 1).contiguous()

torch.testing.assert_close(layer.sglang_w13_weight(), expected_w13)
torch.testing.assert_close(layer.sglang_w2_weight(), expected_w2)


def test_sglang_grouped_mlp_init_does_not_require_grouped_gemm(monkeypatch):
from megatron.core.transformer.moe import grouped_gemm_util as gg

def fail_grouped_gemm_check():
raise AssertionError("grouped_gemm should not be required for SGLangGroupedMLP")

monkeypatch.setattr(gg, "assert_grouped_gemm_is_available", fail_grouped_gemm_check)
config = _make_config(
hidden_size=4,
moe_ffn_hidden_size=3,
num_moe_experts=8,
add_bias_linear=False,
gated_linear_unit=True,
expert_model_parallel_size=4,
)
pg_collection = types.SimpleNamespace(
ep=_FakeCPGroup(4, 1), expt_tp=_FakeCPGroup(1, 0), expt_dp=_FakeCPGroup(1, 0)
)

layer = SGLangGroupedMLP(2, config, pg_collection)

assert layer.weight1.shape == (4, 12)
assert layer.weight2.shape == (6, 4)


def test_sglang_grouped_mlp_masks_nonlocal_topk_ids_for_standard_ep_path():
from megatron.core.transformer.moe.sgl_fused_moe.forward import remap_global_to_local_expert_ids

global_topk_ids = torch.tensor([[0, 3, 2], [1, 2, 3]], dtype=torch.long)

actual = remap_global_to_local_expert_ids(
global_topk_ids, num_experts=4, num_local_experts=2, ep_rank=1, ep_size=2
)

torch.testing.assert_close(actual, torch.tensor([[-1, 1, 0], [-1, 0, 1]], dtype=torch.int32))


def test_sglang_moe_forward_uses_hidden_size_from_sglang_weight_layout(monkeypatch):
from megatron.core.transformer.moe.sgl_fused_moe import forward as moe_forward

captured = {}

def fake_fused_experts(*, hidden_states, w1, w2, topk_output, moe_runner_config):
captured["hidden_states"] = hidden_states
captured["topk_output"] = topk_output
captured["moe_runner_config"] = moe_runner_config
return hidden_states

monkeypatch.setattr(moe_forward, "HAVE_SGLANG_FUSED_MOE", True)
monkeypatch.setattr(moe_forward, "get_global_server_args", lambda: object())
monkeypatch.setattr(
moe_forward, "StandardTopKOutput", lambda **kwargs: types.SimpleNamespace(**kwargs)
)
monkeypatch.setattr(
moe_forward, "MoeRunnerConfig", lambda **kwargs: types.SimpleNamespace(**kwargs)
)
monkeypatch.setattr(moe_forward, "fused_experts", fake_fused_experts)

hidden_states = torch.randn(4, 2, dtype=torch.bfloat16)
w1 = torch.randn(2, 6, 2, dtype=torch.bfloat16)
w2 = torch.randn(2, 2, 3, dtype=torch.bfloat16)
topk_weights = torch.ones(4, 2, dtype=torch.float32)
topk_ids = torch.tensor([[0, 2], [1, 3], [2, 0], [3, 1]], dtype=torch.long)

actual = moe_forward.sglang_moe_forward(
hidden_states,
w1,
w2,
topk_weights,
topk_ids,
num_experts=4,
num_local_experts=2,
ep_rank=1,
ep_size=2,
layer_id=7,
)

torch.testing.assert_close(actual, hidden_states)
runner_config = captured["moe_runner_config"]
assert runner_config.hidden_size == 2
assert runner_config.intermediate_size_per_partition == 3
assert runner_config.layer_id == 7
torch.testing.assert_close(
captured["topk_output"].topk_ids,
torch.tensor([[-1, 0], [-1, 1], [0, -1], [1, -1]], dtype=torch.int32),
)


def test_sglang_mtp_spec_uses_sglang_backend():
config = _make_config(true_on_policy_contract=QWEN3_DENSE_TRUE_ON_POLICY_V1, mtp_num_layers=1)
transformer_layer_spec = get_gpt_decoder_layer_specs(
Expand Down