[MLX] Add correctness tests for qwen2_moe and qwen3_moe (#29440)
This commit is contained in:
@@ -0,0 +1,122 @@
|
|||||||
|
import importlib.util
|
||||||
|
import os
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
popen_launch_server,
|
||||||
|
try_cached_model,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Registered on the CPU suite but skipped wherever mlx is absent; runs for real
|
||||||
|
# only on Apple Silicon. The macOS CI lane (pr-test-mlx.yml) is model-free, so
|
||||||
|
# this serving test is not wired into it and still runs only locally.
|
||||||
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
_HAS_MLX = importlib.util.find_spec("mlx") is not None
|
||||||
|
|
||||||
|
# qwen2_moe architecture (Qwen2MoeForCausalLM), served on the MLX backend.
|
||||||
|
# The model runs through mlx_lm's own qwen2_moe implementation; the SGLang MLX
|
||||||
|
# backend does not require any srt/models file for it. This test is a black-box
|
||||||
|
# correctness guard for the served model.
|
||||||
|
#
|
||||||
|
# Default is the MLX-community 4-bit repo so the test is portable. Override with
|
||||||
|
# SGLANG_MLX_TEST_MODEL to point at a local copy, e.g.
|
||||||
|
# SGLANG_MLX_TEST_MODEL=models/Qwen1.5-MoE-A2.7B-Chat-4bit
|
||||||
|
MODEL_PATH = os.environ.get(
|
||||||
|
"SGLANG_MLX_TEST_MODEL", "mlx-community/Qwen1.5-MoE-A2.7B-Chat-4bit"
|
||||||
|
)
|
||||||
|
|
||||||
|
# mem-fraction is tuned conservatively for a 24 GB Apple Silicon machine.
|
||||||
|
MEM_FRACTION_STATIC = os.environ.get("SGLANG_MLX_TEST_MEM_FRACTION", "0.7")
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_HAS_MLX, "requires mlx (Apple Silicon only)")
|
||||||
|
class TestQwen2MoeMlxCorrectness(CustomTestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = try_cached_model(MODEL_PATH)
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
|
||||||
|
env = os.environ.copy()
|
||||||
|
env["SGLANG_USE_MLX"] = "1"
|
||||||
|
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--tp-size",
|
||||||
|
"1",
|
||||||
|
"--disable-radix-cache",
|
||||||
|
"--disable-cuda-graph",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
MEM_FRACTION_STATIC,
|
||||||
|
"--max-running-requests",
|
||||||
|
"1",
|
||||||
|
"--context-length",
|
||||||
|
"2048",
|
||||||
|
],
|
||||||
|
env=env,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if hasattr(cls, "process") and cls.process is not None:
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
def _chat(self, messages, max_tokens=32, temperature=0):
|
||||||
|
resp = requests.post(
|
||||||
|
f"{self.base_url}/v1/chat/completions",
|
||||||
|
json={
|
||||||
|
"model": MODEL_PATH,
|
||||||
|
"messages": messages,
|
||||||
|
"temperature": temperature,
|
||||||
|
"max_tokens": max_tokens,
|
||||||
|
},
|
||||||
|
timeout=120,
|
||||||
|
)
|
||||||
|
resp.raise_for_status()
|
||||||
|
return resp.json()["choices"][0]["message"]["content"].strip()
|
||||||
|
|
||||||
|
def test_basic_generation_nonempty(self):
|
||||||
|
text = self._chat(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "You are a concise assistant."},
|
||||||
|
{"role": "user", "content": "Say hello briefly."},
|
||||||
|
],
|
||||||
|
max_tokens=16,
|
||||||
|
)
|
||||||
|
self.assertIsInstance(text, str)
|
||||||
|
self.assertGreater(len(text), 0)
|
||||||
|
|
||||||
|
def test_simple_arithmetic(self):
|
||||||
|
text = self._chat(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "You are a concise assistant."},
|
||||||
|
{"role": "user", "content": "What is 2+2? Reply with just the number."},
|
||||||
|
],
|
||||||
|
max_tokens=8,
|
||||||
|
)
|
||||||
|
self.assertIn("4", text)
|
||||||
|
|
||||||
|
def test_simple_fact(self):
|
||||||
|
text = self._chat(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "You are a concise assistant."},
|
||||||
|
{"role": "user", "content": "What is the capital of France? One word."},
|
||||||
|
],
|
||||||
|
max_tokens=8,
|
||||||
|
)
|
||||||
|
self.assertIn("Paris", text)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,122 @@
|
|||||||
|
import importlib.util
|
||||||
|
import os
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
popen_launch_server,
|
||||||
|
try_cached_model,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Registered on the CPU suite but skipped wherever mlx is absent; runs for real
|
||||||
|
# only on Apple Silicon. The macOS CI lane (pr-test-mlx.yml) is model-free, so
|
||||||
|
# this serving test is not wired into it and still runs only locally.
|
||||||
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
_HAS_MLX = importlib.util.find_spec("mlx") is not None
|
||||||
|
|
||||||
|
# qwen3_moe architecture (Qwen3MoeForCausalLM), served on the MLX backend.
|
||||||
|
# The model runs through mlx_lm's own qwen3_moe implementation; the SGLang MLX
|
||||||
|
# backend does not require any srt/models file for it. This test is a black-box
|
||||||
|
# correctness guard for the served model. Qwen3 is a hybrid-thinking model, so
|
||||||
|
# thinking is disabled to keep outputs short and deterministic.
|
||||||
|
#
|
||||||
|
# Default is the MLX-community 4-bit repo so the test is portable. Override with
|
||||||
|
# SGLANG_MLX_TEST_MODEL to point at a local copy, e.g.
|
||||||
|
# SGLANG_MLX_TEST_MODEL=models/Qwen3-30B-A3B-4bit
|
||||||
|
MODEL_PATH = os.environ.get("SGLANG_MLX_TEST_MODEL", "mlx-community/Qwen3-30B-A3B-4bit")
|
||||||
|
|
||||||
|
# mem-fraction is tuned conservatively for a 24 GB Apple Silicon machine.
|
||||||
|
MEM_FRACTION_STATIC = os.environ.get("SGLANG_MLX_TEST_MEM_FRACTION", "0.9")
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_HAS_MLX, "requires mlx (Apple Silicon only)")
|
||||||
|
class TestQwen3MoeMlxCorrectness(CustomTestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = try_cached_model(MODEL_PATH)
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
|
||||||
|
env = os.environ.copy()
|
||||||
|
env["SGLANG_USE_MLX"] = "1"
|
||||||
|
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--tp-size",
|
||||||
|
"1",
|
||||||
|
"--disable-radix-cache",
|
||||||
|
"--disable-cuda-graph",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
MEM_FRACTION_STATIC,
|
||||||
|
"--max-running-requests",
|
||||||
|
"1",
|
||||||
|
"--context-length",
|
||||||
|
"2048",
|
||||||
|
],
|
||||||
|
env=env,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if hasattr(cls, "process") and cls.process is not None:
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
def _chat(self, messages, max_tokens=32, temperature=0):
|
||||||
|
resp = requests.post(
|
||||||
|
f"{self.base_url}/v1/chat/completions",
|
||||||
|
json={
|
||||||
|
"model": MODEL_PATH,
|
||||||
|
"messages": messages,
|
||||||
|
"temperature": temperature,
|
||||||
|
"max_tokens": max_tokens,
|
||||||
|
"chat_template_kwargs": {"enable_thinking": False},
|
||||||
|
},
|
||||||
|
timeout=120,
|
||||||
|
)
|
||||||
|
resp.raise_for_status()
|
||||||
|
return resp.json()["choices"][0]["message"]["content"].strip()
|
||||||
|
|
||||||
|
def test_basic_generation_nonempty(self):
|
||||||
|
text = self._chat(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "You are a concise assistant."},
|
||||||
|
{"role": "user", "content": "Say hello briefly."},
|
||||||
|
],
|
||||||
|
max_tokens=16,
|
||||||
|
)
|
||||||
|
self.assertIsInstance(text, str)
|
||||||
|
self.assertGreater(len(text), 0)
|
||||||
|
|
||||||
|
def test_simple_arithmetic(self):
|
||||||
|
text = self._chat(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "You are a concise assistant."},
|
||||||
|
{"role": "user", "content": "What is 2+2? Reply with just the number."},
|
||||||
|
],
|
||||||
|
max_tokens=8,
|
||||||
|
)
|
||||||
|
self.assertIn("4", text)
|
||||||
|
|
||||||
|
def test_simple_fact(self):
|
||||||
|
text = self._chat(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "You are a concise assistant."},
|
||||||
|
{"role": "user", "content": "What is the capital of France? One word."},
|
||||||
|
],
|
||||||
|
max_tokens=8,
|
||||||
|
)
|
||||||
|
self.assertIn("Paris", text)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,255 @@
|
|||||||
|
"""Reference-equivalence correctness for the MLX MoE execution path.
|
||||||
|
|
||||||
|
The existing ``test/registered/models/test_qwen{2,3}_moe_mlx_correctness.py`` are
|
||||||
|
black-box smoke tests: they only check that a served model emits *plausible* text
|
||||||
|
("Paris" appears, "4" appears). A subtly broken KV cache, slot allocator, or
|
||||||
|
prefill-chunking change can still pass those.
|
||||||
|
|
||||||
|
This is the strong guard. The SGLang MLX backend executes models by wrapping
|
||||||
|
``mlx_lm`` (``mlx_lm.load`` + the model's own forward, with attention patched for
|
||||||
|
SGLang's cache) and decodes greedily (``mx.argmax``). So the ground truth for "did
|
||||||
|
SGLang corrupt the output?" is raw, *unpatched* ``mlx_lm`` greedy generation on the
|
||||||
|
same prompt. We record those reference tokens, then drive ``MlxModelRunner``
|
||||||
|
(prefill + decode_batch) and assert the tokens match exactly, up to and including
|
||||||
|
EOS. Empirically the agreement is exact (not approximate), so any divergence here
|
||||||
|
is a real regression.
|
||||||
|
|
||||||
|
A second test pins batching isolation: a prompt decoded inside a multi-request
|
||||||
|
``decode_batch`` must yield the same tokens as when decoded alone -- guarding the
|
||||||
|
slot/cache bookkeeping that single-request black-box tests never touch.
|
||||||
|
|
||||||
|
Memory safety: the runner patches the model in-place, so the reference needs its
|
||||||
|
own *separate* ``mlx_lm.load``. Loading both copies at once doubles resident weight
|
||||||
|
memory (~2x model) and can trigger an unrecoverable Metal command-buffer OOM that
|
||||||
|
hard-reboots a memory-constrained Mac. So we load the reference, record tokens,
|
||||||
|
fully release it (``del`` + ``gc`` + ``mx.clear_cache``), and only then build the
|
||||||
|
runner -- keeping peak at a single model copy. A pre-flight free-memory check skips
|
||||||
|
(never crashes) when there isn't headroom for even one copy.
|
||||||
|
|
||||||
|
MLX-gated like its siblings: registered on the CPU suite but skipped wherever the
|
||||||
|
``mlx`` package is absent (all current CI runners), so it runs for real only on
|
||||||
|
Apple Silicon. Override the model with ``SGLANG_MLX_TEST_MODEL`` (e.g. a local path
|
||||||
|
or ``mlx-community/Qwen3-30B-A3B-4bit`` for qwen3_moe); bump
|
||||||
|
``SGLANG_MLX_TEST_MIN_FREE_GB`` accordingly for larger models.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import gc
|
||||||
|
import importlib.util
|
||||||
|
import os
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
_HAS_MLX = (
|
||||||
|
importlib.util.find_spec("mlx") is not None
|
||||||
|
and importlib.util.find_spec("mlx_lm") is not None
|
||||||
|
)
|
||||||
|
_SKIP_REASON = "requires mlx + mlx_lm (Apple Silicon only)"
|
||||||
|
|
||||||
|
# Default to the portable mlx-community 4-bit qwen2_moe repo; override to a local
|
||||||
|
# dir or another MoE arch (e.g. Qwen3-30B-A3B-4bit) via the env var.
|
||||||
|
MODEL_PATH = os.environ.get(
|
||||||
|
"SGLANG_MLX_TEST_MODEL", "mlx-community/Qwen1.5-MoE-A2.7B-Chat-4bit"
|
||||||
|
)
|
||||||
|
MEM_FRACTION_STATIC = float(os.environ.get("SGLANG_MLX_TEST_MEM_FRACTION", "0.7"))
|
||||||
|
# Skip (do NOT crash) unless this much system memory is free: enough for ONE
|
||||||
|
# resident model copy plus activations. An MLX Metal OOM is uncatchable and can
|
||||||
|
# reboot the machine, so this guard must fail safe. ~12 GB suits the default
|
||||||
|
# 14.3B-param 4-bit qwen2_moe; raise it for qwen3_moe (~18+).
|
||||||
|
MIN_FREE_GB = float(os.environ.get("SGLANG_MLX_TEST_MIN_FREE_GB", "12"))
|
||||||
|
|
||||||
|
# Short prompts with deterministic, quickly-terminating greedy answers.
|
||||||
|
PROMPTS = [
|
||||||
|
"List the first 10 prime numbers, comma separated.",
|
||||||
|
"What is the capital of France? Answer in one word.",
|
||||||
|
"Write one short sentence about the ocean.",
|
||||||
|
]
|
||||||
|
MAX_NEW_TOKENS = 64 # hard cap; most prompts hit EOS well before this
|
||||||
|
BATCH_HORIZON = 24 # fixed step count for the batching-isolation test
|
||||||
|
|
||||||
|
|
||||||
|
def _available_gb():
|
||||||
|
try:
|
||||||
|
import psutil
|
||||||
|
|
||||||
|
return psutil.virtual_memory().available / 1024**3
|
||||||
|
except Exception:
|
||||||
|
return None # psutil absent -> skip the pre-flight check
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestMlxReferenceCorrectness(CustomTestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
import mlx.core as mx
|
||||||
|
from mlx_lm import load
|
||||||
|
|
||||||
|
avail = _available_gb()
|
||||||
|
if avail is not None and avail < MIN_FREE_GB:
|
||||||
|
raise unittest.SkipTest(
|
||||||
|
f"insufficient free memory: {avail:.1f} GB < {MIN_FREE_GB} GB needed "
|
||||||
|
f"to safely load {MODEL_PATH} (override SGLANG_MLX_TEST_MIN_FREE_GB)"
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- Phase 1: reference tokens from UNPATCHED mlx_lm (one copy resident) ---
|
||||||
|
try:
|
||||||
|
ref_model, cls.tokenizer = load(
|
||||||
|
MODEL_PATH, tokenizer_config={"trust_remote_code": True}
|
||||||
|
)
|
||||||
|
except Exception as exc: # not cached / offline / bad path
|
||||||
|
raise unittest.SkipTest(f"could not load {MODEL_PATH}: {exc}")
|
||||||
|
|
||||||
|
eos = getattr(cls.tokenizer, "eos_token_ids", None) or {
|
||||||
|
cls.tokenizer.eos_token_id
|
||||||
|
}
|
||||||
|
cls.eos_ids = set(eos)
|
||||||
|
|
||||||
|
cls.cases = [] # (prompt, prompt_ids, reference_token_ids)
|
||||||
|
for prompt in PROMPTS:
|
||||||
|
prompt_ids = list(
|
||||||
|
cls.tokenizer.apply_chat_template(
|
||||||
|
[{"role": "user", "content": prompt}], add_generation_prompt=True
|
||||||
|
)
|
||||||
|
)
|
||||||
|
ref_ids = cls._reference_greedy(
|
||||||
|
ref_model, cls.tokenizer, prompt_ids, MAX_NEW_TOKENS
|
||||||
|
)
|
||||||
|
cls.cases.append((prompt, prompt_ids, ref_ids))
|
||||||
|
|
||||||
|
# --- Release the reference BEFORE building the runner (cap peak at 1x) ---
|
||||||
|
del ref_model
|
||||||
|
gc.collect()
|
||||||
|
mx.clear_cache()
|
||||||
|
active_gb = mx.get_active_memory() / 1024**3
|
||||||
|
if active_gb > 2.0:
|
||||||
|
# Weights were not released; loading the runner now would double resident
|
||||||
|
# memory and risk an OOM. Bail safely instead.
|
||||||
|
raise unittest.SkipTest(
|
||||||
|
f"reference model not released (active={active_gb:.1f} GB); "
|
||||||
|
"skipping to avoid a double-resident OOM"
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- Phase 2: SGLang runner (one copy resident) ---
|
||||||
|
from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner
|
||||||
|
|
||||||
|
cls.runner = MlxModelRunner(
|
||||||
|
model_path=MODEL_PATH,
|
||||||
|
trust_remote_code=True,
|
||||||
|
disable_radix_cache=True, # per-request contiguous caches; no big pool
|
||||||
|
mem_fraction_static=MEM_FRACTION_STATIC,
|
||||||
|
)
|
||||||
|
cls.runner.init_cache_pools(req_to_token_pool=None)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
runner = getattr(cls, "runner", None)
|
||||||
|
if runner is not None:
|
||||||
|
runner.clear()
|
||||||
|
cls.runner = None
|
||||||
|
gc.collect()
|
||||||
|
try:
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
mx.clear_cache()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# --- helpers ----------------------------------------------------------
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _reference_greedy(model, tokenizer, prompt_ids, max_new):
|
||||||
|
"""Ground-truth token ids from raw, unpatched mlx_lm greedy generation."""
|
||||||
|
import mlx.core as mx
|
||||||
|
from mlx_lm import stream_generate
|
||||||
|
from mlx_lm.sample_utils import make_sampler
|
||||||
|
|
||||||
|
sampler = make_sampler(temp=0.0) # greedy / argmax
|
||||||
|
out = []
|
||||||
|
for resp in stream_generate(
|
||||||
|
model, tokenizer, mx.array(prompt_ids), max_tokens=max_new, sampler=sampler
|
||||||
|
):
|
||||||
|
out.append(int(resp.token))
|
||||||
|
return out
|
||||||
|
|
||||||
|
def _prefill(self, rid, prompt_ids):
|
||||||
|
return int(
|
||||||
|
self.runner.prefill(
|
||||||
|
req_id=rid,
|
||||||
|
new_token_ids=list(prompt_ids),
|
||||||
|
full_token_ids=list(prompt_ids),
|
||||||
|
prefix_slot_ids=[],
|
||||||
|
new_slot_ids=[],
|
||||||
|
req_pool_idx=0,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def _decode(self, rids):
|
||||||
|
return [int(t) for t in self.runner.decode_batch(rids)]
|
||||||
|
|
||||||
|
def _sglang_greedy(self, rid, prompt_ids, max_new):
|
||||||
|
"""SGLang MLX greedy generation, stopping at EOS like the reference."""
|
||||||
|
tok = self._prefill(rid, prompt_ids)
|
||||||
|
out = [tok]
|
||||||
|
while len(out) < max_new and tok not in self.eos_ids:
|
||||||
|
tok = self._decode([rid])[0]
|
||||||
|
out.append(tok)
|
||||||
|
self.runner.remove_request(rid)
|
||||||
|
return out
|
||||||
|
|
||||||
|
def _diff_msg(self, prompt, ref, sgl):
|
||||||
|
horizon = min(len(ref), len(sgl))
|
||||||
|
first = next((j for j in range(horizon) if ref[j] != sgl[j]), horizon)
|
||||||
|
return (
|
||||||
|
f"\nprompt: {prompt!r}"
|
||||||
|
f"\n first divergence @ index {first} (len ref={len(ref)} sgl={len(sgl)})"
|
||||||
|
f"\n ref text: {self.tokenizer.decode(ref)!r}"
|
||||||
|
f"\n sgl text: {self.tokenizer.decode(sgl)!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- tests ------------------------------------------------------------
|
||||||
|
|
||||||
|
def test_greedy_matches_reference_exact(self):
|
||||||
|
"""SGLang MLX greedy output == unpatched mlx_lm greedy output, token-for-token."""
|
||||||
|
for i, (prompt, prompt_ids, ref) in enumerate(self.cases):
|
||||||
|
sgl = self._sglang_greedy(f"ref-{i}", prompt_ids, MAX_NEW_TOKENS)
|
||||||
|
self.assertEqual(sgl, ref, self._diff_msg(prompt, ref, sgl))
|
||||||
|
|
||||||
|
def test_batched_decode_matches_solo(self):
|
||||||
|
"""A request's tokens are identical whether decoded alone or in a batch.
|
||||||
|
|
||||||
|
Pins slot/cache isolation: ``decode_batch`` over several concurrent
|
||||||
|
requests must not let one request's state bleed into another's.
|
||||||
|
"""
|
||||||
|
ids_list = [prompt_ids for (_, prompt_ids, _) in self.cases]
|
||||||
|
|
||||||
|
# Solo: prefill, decode a fixed horizon, remove -- one request at a time.
|
||||||
|
solo = []
|
||||||
|
for i, ids in enumerate(ids_list):
|
||||||
|
seq = [self._prefill(f"solo-{i}", ids)]
|
||||||
|
for _ in range(BATCH_HORIZON - 1):
|
||||||
|
seq.append(self._decode([f"solo-{i}"])[0])
|
||||||
|
self.runner.remove_request(f"solo-{i}")
|
||||||
|
solo.append(seq)
|
||||||
|
|
||||||
|
# Batched: prefill all, then advance them together in one decode_batch.
|
||||||
|
rids = [f"batch-{i}" for i in range(len(ids_list))]
|
||||||
|
batched = [[self._prefill(rid, ids)] for rid, ids in zip(rids, ids_list)]
|
||||||
|
for _ in range(BATCH_HORIZON - 1):
|
||||||
|
for j, t in enumerate(self._decode(rids)):
|
||||||
|
batched[j].append(t)
|
||||||
|
for rid in rids:
|
||||||
|
self.runner.remove_request(rid)
|
||||||
|
|
||||||
|
for i, (prompt, _, _) in enumerate(self.cases):
|
||||||
|
self.assertEqual(
|
||||||
|
batched[i], solo[i], self._diff_msg(prompt, solo[i], batched[i])
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main(verbosity=3)
|
||||||
Reference in New Issue
Block a user