Support for Zyphra zaya1 model (#26347)

Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
This commit is contained in:
ChengYao-amd
2026-06-10 02:44:47 -07:00
committed by GitHub
co-authored by github-actions[bot]
parent b40f365732
commit 255843d454
9 changed files with 3167 additions and 1 deletions
+127
View File
@@ -0,0 +1,127 @@
"""End-to-end server test for Zyphra ZAYA1 (hybrid CCA attention + MoE).
This test boots a real ``Zyphra/ZAYA1-base`` SGLang server via
``popen_launch_server``, sends a handful of completions through the HTTP API,
and finishes with a small MMLU sanity slice.
The test is gated behind ``RUN_ZAYA_E2E=1`` so the registered suite does not
have to download the full ZAYA1-base checkpoint (≈17 GB) on every run; the CI
job that owns this test sets the variable explicitly.
"""
import os
import unittest
from types import SimpleNamespace
from sglang.srt.utils import is_hip, kill_process_tree
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.run_eval import run_eval
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
# ZAYA1-base is a heavyweight launch (≈120 transformer layers with MoE), so
# the estimated time is set generously to keep the CI scheduler from preempting
# the job before the server finishes warming up.
register_cuda_ci(est_time=420, stage="extra-a", runner_config="1-gpu-large")
register_amd_ci(est_time=420, suite="stage-b-test-1-gpu-large-amd")
_MODEL_PATH = os.environ.get("ZAYA_MODEL_PATH", "Zyphra/ZAYA1-base")
def _zaya_enabled() -> bool:
return os.environ.get("RUN_ZAYA_E2E", "0") == "1"
@unittest.skipUnless(
_zaya_enabled(),
"Set RUN_ZAYA_E2E=1 to enable the ZAYA1 end-to-end server test "
"(requires downloading the model weights).",
)
class TestZayaServer(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = _MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
other_args = [
"--mem-fraction-static",
"0.5",
"--max-running-requests",
"8",
]
if is_hip():
other_args += ["--attention-backend", "triton"]
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=other_args,
)
@classmethod
def tearDownClass(cls):
if getattr(cls, "process", None) is not None:
kill_process_tree(cls.process.pid)
def test_generation_basic(self):
"""Send three prompts through the ``/generate`` endpoint and require
non-empty completions for each."""
import requests
prompts = [
"The capital of France is",
"1 + 2 + 3 + 4 + 5 =",
"Write a haiku about silicon:",
]
for prompt in prompts:
resp = requests.post(
f"{self.base_url}/generate",
json={
"text": prompt,
"sampling_params": {
"temperature": 0.0,
"max_new_tokens": 16,
},
},
timeout=60,
)
self.assertEqual(resp.status_code, 200, resp.text)
data = resp.json()
self.assertIn("text", data, data)
self.assertGreater(len(data["text"].strip()), 0, data)
def test_mmlu_sanity(self):
"""32-example MMLU sanity slice.
ZAYA1-base is a pretrained (non instruction-tuned) checkpoint that
emits long ``<think>…</think>`` reasoning blocks before settling on a
final letter, so ``max_tokens`` must be large enough for the evaluator
to see the chosen answer. The threshold sits just above chance: it is
a regression sanity check rather than a production-quality gate. An
instruction-tuned ZAYA1 checkpoint scores meaningfully higher and
should raise this bound when wired in.
"""
args = SimpleNamespace(
base_url=self.base_url,
model=self.model,
eval_name="mmlu",
num_examples=32,
num_threads=8,
max_tokens=1024,
)
metrics = run_eval(args)
self.assertGreaterEqual(
metrics["score"],
0.30,
f"MMLU sanity below threshold: {metrics}",
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,101 @@
"""Unit tests for ``sglang.srt.configs.zaya.ZayaConfig``."""
import unittest
from transformers import AutoConfig
from sglang.srt.configs.zaya import ZayaConfig, register_zaya_config
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class TestZayaConfig(CustomTestCase):
def test_default_fields_match_zaya1_base(self):
"""Defaults reflect Zyphra/ZAYA1-base reference checkpoint."""
cfg = ZayaConfig()
self.assertEqual(cfg.model_type, "zaya")
self.assertEqual(cfg.hidden_size, 2048)
self.assertEqual(cfg.head_dim, 128)
self.assertEqual(cfg.num_attention_heads, 8)
self.assertEqual(cfg.num_query_groups, 2)
self.assertEqual(cfg.num_key_value_heads, 2)
self.assertEqual(cfg.num_experts, 16)
self.assertEqual(cfg.moe_router_topk, 1)
self.assertEqual(cfg.ffn_hidden_size, 4096)
self.assertEqual(cfg.zaya_mlp_expansion, 256)
self.assertEqual(cfg.cca_time0, 2)
self.assertEqual(cfg.cca_time1, 2)
self.assertTrue(cfg.tie_word_embeddings)
self.assertTrue(cfg.zaya_use_eda)
self.assertTrue(cfg.zaya_use_mod)
self.assertTrue(cfg.scale_residual_merge)
self.assertEqual(cfg.partial_rotary_factor, 0.5)
self.assertEqual(cfg.rope_theta, 1_000_000.0)
def test_rope_parameters_auto_derived(self):
"""When neither ``rope_scaling`` nor ``rope_parameters`` is supplied,
both ``rope_theta`` and ``partial_rotary_factor`` should still appear
inside ``rope_parameters`` together with a default ``rope_type``.
"""
cfg = ZayaConfig()
rp = cfg.rope_parameters
self.assertEqual(rp["rope_type"], "default")
self.assertEqual(rp["rope_theta"], 1_000_000.0)
self.assertEqual(rp["partial_rotary_factor"], 0.5)
def test_rope_parameters_explicit_takes_priority(self):
cfg = ZayaConfig(rope_parameters={"type": "linear", "factor": 4.0})
rp = cfg.rope_parameters
# ``type`` is normalized to ``rope_type``.
self.assertEqual(rp["rope_type"], "linear")
self.assertEqual(rp["factor"], 4.0)
# Defaults are still merged in.
self.assertEqual(rp["rope_theta"], 1_000_000.0)
def test_head_dim_required(self):
with self.assertRaises(AssertionError):
ZayaConfig(head_dim=None)
def test_num_query_groups_must_equal_kv_heads(self):
with self.assertRaises(AssertionError):
ZayaConfig(num_query_groups=4, num_key_value_heads=2)
def test_hybrid_model_properties(self):
"""Verify properties required for HybridReqToTokenPool integration."""
cfg = ZayaConfig()
# Default 80 layers: even layers are attention, odd are MoE
self.assertEqual(cfg.full_attention_layer_ids, list(range(0, 80, 2)))
self.assertEqual(cfg.linear_layer_ids, cfg.full_attention_layer_ids)
self.assertEqual(cfg.mamba_chunk_size, 1)
params = cfg.mamba2_cache_params
self.assertIsNotNone(params)
# conv[0] = conv_state: (in_out_ch, total_padding)
in_out_ch = (cfg.num_attention_heads + cfg.num_key_value_heads) * cfg.head_dim
total_padding = (cfg.cca_time0 - 1) + (cfg.cca_time1 - 1)
self.assertEqual(params.shape.conv[0], (in_out_ch, total_padding))
# conv[1] = prev_hs: (hidden_size, 1)
self.assertEqual(params.shape.conv[1], (cfg.hidden_size, 1))
self.assertEqual(params.layers, cfg.linear_layer_ids)
def test_hybrid_model_properties_with_zaya_layers(self):
"""When zaya_layers is provided, layer IDs derive from the list."""
cfg = ZayaConfig(zaya_layers=["a", 16, "a", 16])
self.assertEqual(cfg.num_hidden_layers, 4)
self.assertEqual(cfg.full_attention_layer_ids, [0, 2])
self.assertEqual(cfg.linear_layer_ids, [0, 2])
def test_auto_config_registration_is_idempotent(self):
# Calling the helper twice must not raise even though importing the
# module already registered the model type.
register_zaya_config()
register_zaya_config()
# ``AutoConfig.for_model`` now resolves to ``ZayaConfig``.
cfg = AutoConfig.for_model("zaya")
self.assertIsInstance(cfg, ZayaConfig)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,765 @@
"""Numerical and state-cache correctness tests for the ZAYA1 CCA module.
The CCA per-request conv-state cache must satisfy the following invariants,
which are each exercised by a dedicated test case:
1. A single-chunk extend forward (no prefix) is numerically equivalent to the
reference torch implementation that processes the whole sequence at once.
2. Splitting a sequence into one prefill of ``S0`` tokens and ``S1`` single-
token decode steps produces the same q / k / v tensors as the equivalent
single-chunk run.
3. A batched two-request decode for request 0 yields identical q / k / v to a
single-request decode of request 0 at the same step.
4. Multi-request prefills update only the conv state and ``prev_hs`` slots for
each request and leave unused slots zero.
5. A simulated tensor-parallel (TP=2) CCA produces per-rank q / k / v slices
that match the corresponding head slices of a TP=1 reference, both for
prefill (``_forward_extend``) and for decode (``_forward_decode``).
All tests run on CPU with a tiny configuration so they stay fast and have no
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
from types import SimpleNamespace
from typing import List, Optional
import torch
from sglang.test.ci.ci_register import register_cpu_ci
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
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class _MockLayerCache:
conv: List[torch.Tensor]
temporal: torch.Tensor
class _MockReqToTokenPool:
"""Minimal stand-in for ``HybridReqToTokenPool`` providing the two methods
that CCA calls: ``mamba2_layer_cache`` and ``get_mamba_indices``.
For TP-aware tests, ``tp_size`` controls the per-rank ``in_out_ch`` of the
``conv[0]`` state. ``conv[1]`` (prev_hs) is replicated and stays at full
``hidden_size``.
"""
def __init__(self, pool_size: int, cca_config, tp_size: int = 1):
in_out_ch_full = (
cca_config.num_attention_heads + cca_config.num_key_value_heads
) * cca_config.head_dim
assert in_out_ch_full % tp_size == 0
in_out_ch_per_rank = in_out_ch_full // tp_size
total_padding = (cca_config.cca_time0 - 1) + (cca_config.cca_time1 - 1)
num_layers = len(cca_config.linear_layer_ids)
self.conv_state = torch.zeros(
num_layers, pool_size + 1, in_out_ch_per_rank, total_padding
)
self.prev_hs_state = torch.zeros(
num_layers, pool_size + 1, cca_config.hidden_size, 1
)
self.temporal = torch.zeros(num_layers, pool_size + 1, 1, 1, 0)
self._layer_map = {lid: i for i, lid in enumerate(cca_config.linear_layer_ids)}
self._identity_map = torch.arange(pool_size + 1, dtype=torch.int32)
def mamba2_layer_cache(self, layer_id: int):
idx = self._layer_map[layer_id]
return _MockLayerCache(
conv=[self.conv_state[idx], self.prev_hs_state[idx]],
temporal=self.temporal[idx],
)
def get_mamba_indices(self, req_pool_indices: torch.Tensor) -> torch.Tensor:
return req_pool_indices.to(torch.int32)
@contextmanager
def _mock_pool_context(pool: _MockReqToTokenPool):
"""Install a mock ``ForwardContext`` whose ``req_to_token_pool`` is ``pool``."""
from sglang.srt.model_executor.forward_context import (
ForwardContext,
set_forward_context,
)
backend = SimpleNamespace(req_to_token_pool=pool, token_to_kv_pool=None)
ctx = ForwardContext(attn_backend=backend)
prev = set_forward_context(ctx)
try:
yield pool
finally:
set_forward_context(prev)
# ---------------------------------------------------------------------------
# Helper factories
# ---------------------------------------------------------------------------
def _make_forward_batch(
*,
is_decode: bool,
extend_seq_lens_cpu,
extend_prefix_lens_cpu,
req_pool_indices,
input_ids: torch.Tensor,
):
from sglang.srt.model_executor.forward_batch_info import ForwardMode
mode = ForwardMode.DECODE if is_decode else ForwardMode.EXTEND
forward_batch = SimpleNamespace()
forward_batch.forward_mode = mode
forward_batch.input_ids = input_ids
forward_batch.req_pool_indices = torch.as_tensor(
req_pool_indices, dtype=torch.int32
)
forward_batch.extend_seq_lens_cpu = list(extend_seq_lens_cpu)
forward_batch.extend_prefix_lens_cpu = list(extend_prefix_lens_cpu)
return forward_batch
def _make_tiny_config(num_hidden_layers: int = 2):
from sglang.srt.configs.zaya import ZayaConfig
return ZayaConfig(
hidden_size=16,
ffn_hidden_size=32,
num_hidden_layers=num_hidden_layers,
num_experts=2,
num_attention_heads=4,
num_query_groups=2,
num_key_value_heads=2,
head_dim=8,
cca_time0=2,
cca_time1=2,
max_position_embeddings=64,
moe_router_topk=1,
zaya_mlp_expansion=8,
attention_bias=False,
)
def _make_tiny_cca(
seed: int = 0,
tp_rank: Optional[int] = None,
tp_size: Optional[int] = None,
layer_id: int = 0,
config=None,
):
from sglang.srt.models.zaya import CCA
if config is None:
config = _make_tiny_config()
torch.manual_seed(seed)
cca = CCA(
config=config,
cca_num_k_heads=config.num_query_groups,
cca_num_q_heads=config.num_attention_heads,
hidden_size=config.hidden_size,
head_dim=config.head_dim,
cca_time0=config.cca_time0,
cca_time1=config.cca_time1,
layer_id=layer_id,
tp_rank=tp_rank,
tp_size=tp_size,
)
cca.eval()
with torch.no_grad():
for p in cca.parameters():
p.data.normal_(mean=0.0, std=0.05)
cca.temp.data.zero_()
return cca, config
class TestZayaCCA(CustomTestCase):
@classmethod
def setUpClass(cls) -> None:
_ensure_dist_initialized()
def test_single_chunk_matches_reference(self):
"""A single-chunk extend with empty prefix matches the no-state path."""
cca, config = _make_tiny_cca(seed=1)
cca_ref, _ = _make_tiny_cca(seed=1)
with torch.no_grad():
cca_ref.load_state_dict(cca.state_dict())
S = 5
hs = torch.randn(S, cca.hidden_size, dtype=torch.float32) * 0.1
q_ref, k_ref, v_ref = cca_ref._forward_no_state(hs)
pool = _MockReqToTokenPool(pool_size=8, cca_config=config)
fb = _make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S],
extend_prefix_lens_cpu=[0],
req_pool_indices=[0],
input_ids=torch.arange(S, dtype=torch.int64),
)
with _mock_pool_context(pool):
q, k, v = cca.forward(hs, fb)
torch.testing.assert_close(q, q_ref, atol=1e-5, rtol=1e-5)
torch.testing.assert_close(k, k_ref, atol=1e-5, rtol=1e-5)
torch.testing.assert_close(v, v_ref, atol=1e-5, rtol=1e-5)
def test_prefill_then_decode_matches_full_sequence(self):
"""Prefill(S0) followed by ``S1`` single-token decode steps matches a
one-shot reference over ``S0 + S1`` tokens."""
cca, config = _make_tiny_cca(seed=2)
cca_ref, _ = _make_tiny_cca(seed=2)
with torch.no_grad():
cca_ref.load_state_dict(cca.state_dict())
S0, S1 = 4, 2
S_total = S0 + S1
torch.manual_seed(77)
hs = torch.randn(S_total, cca.hidden_size, dtype=torch.float32) * 0.1
q_ref, k_ref, v_ref = cca_ref._forward_no_state(hs)
pool = _MockReqToTokenPool(pool_size=8, cca_config=config)
with _mock_pool_context(pool):
fb_prefill = _make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S0],
extend_prefix_lens_cpu=[0],
req_pool_indices=[0],
input_ids=torch.arange(S0, dtype=torch.int64),
)
q0, k0, v0 = cca.forward(hs[:S0], fb_prefill)
q_decodes = [q0]
k_decodes = [k0]
v_decodes = [v0]
for t in range(S1):
fb_decode = _make_forward_batch(
is_decode=True,
extend_seq_lens_cpu=[],
extend_prefix_lens_cpu=[],
req_pool_indices=[0],
input_ids=torch.tensor([0], dtype=torch.int64),
)
qd, kd, vd = cca.forward(hs[S0 + t : S0 + t + 1], fb_decode)
q_decodes.append(qd)
k_decodes.append(kd)
v_decodes.append(vd)
q_cat = torch.cat(q_decodes, dim=0)
k_cat = torch.cat(k_decodes, dim=0)
v_cat = torch.cat(v_decodes, dim=0)
torch.testing.assert_close(q_cat, q_ref, atol=1e-4, rtol=1e-4)
torch.testing.assert_close(k_cat, k_ref, atol=1e-4, rtol=1e-4)
torch.testing.assert_close(v_cat, v_ref, atol=1e-4, rtol=1e-4)
def test_batched_decode_matches_single_decode(self):
"""A two-request batched decode of request 0 must produce the same
q / k / v tensors as a single-request decode of request 0."""
cca_single, config = _make_tiny_cca(seed=11)
cca_batched, _ = _make_tiny_cca(seed=11)
with torch.no_grad():
cca_batched.load_state_dict(cca_single.state_dict())
S0 = 4
torch.manual_seed(202)
hs0 = torch.randn(S0, cca_single.hidden_size, dtype=torch.float32) * 0.1
hs1 = torch.randn(S0, cca_single.hidden_size, dtype=torch.float32) * 0.1
decode0 = torch.randn(cca_single.hidden_size, dtype=torch.float32) * 0.1
decode1 = torch.randn(cca_single.hidden_size, dtype=torch.float32) * 0.1
pool_single = _MockReqToTokenPool(pool_size=8, cca_config=config)
with _mock_pool_context(pool_single):
cca_single.forward(
hs0,
_make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S0],
extend_prefix_lens_cpu=[0],
req_pool_indices=[0],
input_ids=torch.arange(S0, dtype=torch.int64),
),
)
q_solo, k_solo, v_solo = cca_single.forward(
decode0.unsqueeze(0),
_make_forward_batch(
is_decode=True,
extend_seq_lens_cpu=[],
extend_prefix_lens_cpu=[],
req_pool_indices=[0],
input_ids=torch.tensor([0], dtype=torch.int64),
),
)
pool_batched = _MockReqToTokenPool(pool_size=8, cca_config=config)
with _mock_pool_context(pool_batched):
cca_batched.forward(
torch.cat([hs0, hs1], dim=0),
_make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S0, S0],
extend_prefix_lens_cpu=[0, 0],
req_pool_indices=[0, 1],
input_ids=torch.arange(2 * S0, dtype=torch.int64),
),
)
q_batch, k_batch, v_batch = cca_batched.forward(
torch.stack([decode0, decode1], dim=0),
_make_forward_batch(
is_decode=True,
extend_seq_lens_cpu=[],
extend_prefix_lens_cpu=[],
req_pool_indices=[0, 1],
input_ids=torch.tensor([0, 1], dtype=torch.int64),
),
)
torch.testing.assert_close(q_batch[0:1], q_solo, atol=1e-5, rtol=1e-5)
torch.testing.assert_close(k_batch[0:1], k_solo, atol=1e-5, rtol=1e-5)
torch.testing.assert_close(v_batch[0:1], v_solo, atol=1e-5, rtol=1e-5)
def test_two_requests_state_isolation(self):
"""A batched prefill of two requests must update only the requests'
own slots in the centralized pool."""
cca, config = _make_tiny_cca(seed=4)
S0, S1 = 3, 2
hs0 = torch.randn(S0, cca.hidden_size, dtype=torch.float32) * 0.1
hs1 = torch.randn(S1, cca.hidden_size, dtype=torch.float32) * 0.1
hs = torch.cat([hs0, hs1], dim=0)
pool = _MockReqToTokenPool(pool_size=8, cca_config=config)
fb = _make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S0, S1],
extend_prefix_lens_cpu=[0, 0],
req_pool_indices=[2, 5],
input_ids=torch.arange(S0 + S1, dtype=torch.int64),
)
with _mock_pool_context(pool):
cca.forward(hs, fb)
layer_cache = pool.mamba2_layer_cache(0)
conv_state = layer_cache.conv[0]
prev_hs_state = layer_cache.conv[1]
self.assertTrue(torch.any(conv_state[2] != 0))
self.assertTrue(torch.any(conv_state[5] != 0))
torch.testing.assert_close(
prev_hs_state[2].squeeze(-1).to(torch.float32),
hs0[-1].to(torch.float32),
atol=1e-5,
rtol=1e-5,
)
torch.testing.assert_close(
prev_hs_state[5].squeeze(-1).to(torch.float32),
hs1[-1].to(torch.float32),
atol=1e-5,
rtol=1e-5,
)
for idx in (0, 1, 3, 4):
self.assertTrue(torch.all(conv_state[idx] == 0))
self.assertTrue(torch.all(prev_hs_state[idx] == 0))
def test_mamba_indices_resolved_once_per_forward_step(self):
"""The req -> MambaPool-slot mapping is identical for every CCA layer in
a step, so it (and its GPU->CPU ``.tolist()`` sync) must be resolved once
per forward step and shared across layers, not recomputed per layer.
Regression guard for the per-layer mamba-sync fix: two CCA layers driven
by a single ForwardBatch must trigger exactly one ``get_mamba_indices``
lookup and one host materialization for the whole step.
"""
class _CountingPool(_MockReqToTokenPool):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.get_mamba_indices_calls = 0
def get_mamba_indices(self, req_pool_indices):
self.get_mamba_indices_calls += 1
return super().get_mamba_indices(req_pool_indices)
# num_hidden_layers=4 -> CCA (even) layers live at ids 0 and 2.
config = _make_tiny_config(num_hidden_layers=4)
self.assertEqual(config.linear_layer_ids, [0, 2])
cca0, _ = _make_tiny_cca(seed=5, layer_id=0, config=config)
cca2, _ = _make_tiny_cca(seed=6, layer_id=2, config=config)
S = 4
hs = torch.randn(S, config.hidden_size, dtype=torch.float32) * 0.1
def _fresh_fb():
return _make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S],
extend_prefix_lens_cpu=[0],
req_pool_indices=[0],
input_ids=torch.arange(S, dtype=torch.int64),
)
pool = _CountingPool(pool_size=8, cca_config=config)
with _mock_pool_context(pool):
fb = _fresh_fb()
cca0.forward(hs, fb)
cca2.forward(hs, fb)
# Two CCA layers, one forward step -> one shared lookup, both the
# device tensor and its host mirror memoized on the ForwardBatch.
self.assertEqual(pool.get_mamba_indices_calls, 1)
self.assertTrue(hasattr(fb, "_zaya_mamba_indices"))
self.assertTrue(hasattr(fb, "_zaya_mamba_indices_cpu"))
self.assertEqual(fb._zaya_mamba_indices_cpu, [0])
# A new forward step (fresh ForwardBatch) resolves the mapping again.
cca0.forward(hs, _fresh_fb())
self.assertEqual(pool.get_mamba_indices_calls, 2)
def test_decode_path_does_not_sync_indices_to_host(self):
"""The decode path indexes the pool entirely on-device, so it must not
populate the host-side index cache (keeping it CUDA-graph friendly)."""
cca, config = _make_tiny_cca(seed=7)
pool = _MockReqToTokenPool(pool_size=8, cca_config=config)
with _mock_pool_context(pool):
cca.forward(
torch.randn(3, config.hidden_size, dtype=torch.float32) * 0.1,
_make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[3],
extend_prefix_lens_cpu=[0],
req_pool_indices=[0],
input_ids=torch.arange(3, dtype=torch.int64),
),
)
fb_decode = _make_forward_batch(
is_decode=True,
extend_seq_lens_cpu=[],
extend_prefix_lens_cpu=[],
req_pool_indices=[0],
input_ids=torch.tensor([0], dtype=torch.int64),
)
cca.forward(
torch.randn(1, config.hidden_size, dtype=torch.float32) * 0.1,
fb_decode,
)
# Device indices are memoized, but the host ``.tolist()`` mirror is only
# built by the extend path.
self.assertTrue(hasattr(fb_decode, "_zaya_mamba_indices"))
self.assertFalse(hasattr(fb_decode, "_zaya_mamba_indices_cpu"))
class TestZayaCCATensorParallel(CustomTestCase):
"""Head-parallel TP equivalence:
For each TP rank, the CCA's q / k / v output must equal the head slice of
the TP=1 reference's output that corresponds to that rank's heads. This
verifies that the grouped-mean step and ``conv_qk.1`` (groups = num_q_heads
+ num_k_heads) are correctly partitioned across heads with no cross-rank
leakage.
"""
TP_SIZE = 2
@classmethod
def setUpClass(cls) -> None:
_ensure_dist_initialized()
def _slice_full_state_dict_into_rank(self, ref_cca, tp_cca, tp_rank: int):
"""Copy the reference's full weights into the per-rank CCA, using the
per-parameter ``weight_loader`` that the CCA installs on its own
parameters during ``__init__``. This mirrors what
``ZayaForCausalLM.load_weights`` does at serving time and is the
only way TP correctness is exercised end-to-end.
"""
ref_state = dict(ref_cca.state_dict())
from sglang.srt.model_loader.weight_utils import default_weight_loader
with torch.no_grad():
for name, param in tp_cca.named_parameters():
full_weight = ref_state[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, full_weight)
def _check_per_rank_outputs(
self,
full_q: torch.Tensor,
full_k: torch.Tensor,
full_v: torch.Tensor,
rank_q: torch.Tensor,
rank_k: torch.Tensor,
rank_v: torch.Tensor,
tp_rank: int,
cfg,
):
"""Compare a TP=2 rank's output against the corresponding head slice
of the TP=1 reference output. Q heads and K heads are partitioned
contiguously across ranks: rank ``r`` owns
``[r*Q_per_rank, (r+1)*Q_per_rank)`` for Q and similarly for K.
"""
q_heads_per_rank = cfg.num_attention_heads // self.TP_SIZE
k_heads_per_rank = cfg.num_query_groups // self.TP_SIZE
q_lo, q_hi = tp_rank * q_heads_per_rank, (tp_rank + 1) * q_heads_per_rank
k_lo, k_hi = tp_rank * k_heads_per_rank, (tp_rank + 1) * k_heads_per_rank
torch.testing.assert_close(
rank_q, full_q[:, q_lo:q_hi, :], atol=1e-5, rtol=1e-5
)
torch.testing.assert_close(
rank_k, full_k[:, k_lo:k_hi, :], atol=1e-5, rtol=1e-5
)
torch.testing.assert_close(
rank_v, full_v[:, k_lo:k_hi, :], atol=1e-5, rtol=1e-5
)
def test_tp2_extend_matches_full(self):
"""Single-chunk extend with TP=2 produces the same q / k / v slices
as a TP=1 reference, verified rank-by-rank.
"""
ref_cca, cfg = _make_tiny_cca(seed=21, tp_rank=0, tp_size=1)
S = 6
torch.manual_seed(901)
hs = torch.randn(S, ref_cca.hidden_size, dtype=torch.float32) * 0.1
ref_pool = _MockReqToTokenPool(pool_size=8, cca_config=cfg, tp_size=1)
ref_fb = _make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S],
extend_prefix_lens_cpu=[0],
req_pool_indices=[0],
input_ids=torch.arange(S, dtype=torch.int64),
)
with _mock_pool_context(ref_pool):
full_q, full_k, full_v = ref_cca.forward(hs, ref_fb)
for tp_rank in range(self.TP_SIZE):
rank_cca, _ = _make_tiny_cca(
seed=21 + tp_rank, tp_rank=tp_rank, tp_size=self.TP_SIZE
)
self._slice_full_state_dict_into_rank(ref_cca, rank_cca, tp_rank)
rank_pool = _MockReqToTokenPool(
pool_size=8, cca_config=cfg, tp_size=self.TP_SIZE
)
rank_fb = _make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S],
extend_prefix_lens_cpu=[0],
req_pool_indices=[0],
input_ids=torch.arange(S, dtype=torch.int64),
)
with _mock_pool_context(rank_pool):
rank_q, rank_k, rank_v = rank_cca.forward(hs, rank_fb)
self._check_per_rank_outputs(
full_q, full_k, full_v, rank_q, rank_k, rank_v, tp_rank, cfg
)
def test_tp2_decode_matches_full(self):
"""Prefill(S0) + decode(1 token) with TP=2 produces the same q / k / v
slices as a TP=1 reference, verifying that the per-rank conv state
and prev_hs cache (which is replicated on every rank) agree.
"""
ref_cca, cfg = _make_tiny_cca(seed=22, tp_rank=0, tp_size=1)
S0 = 5
torch.manual_seed(902)
hs_prefill = torch.randn(S0, ref_cca.hidden_size, dtype=torch.float32) * 0.1
hs_decode = torch.randn(1, ref_cca.hidden_size, dtype=torch.float32) * 0.1
ref_pool = _MockReqToTokenPool(pool_size=8, cca_config=cfg, tp_size=1)
with _mock_pool_context(ref_pool):
ref_cca.forward(
hs_prefill,
_make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S0],
extend_prefix_lens_cpu=[0],
req_pool_indices=[0],
input_ids=torch.arange(S0, dtype=torch.int64),
),
)
full_q, full_k, full_v = ref_cca.forward(
hs_decode,
_make_forward_batch(
is_decode=True,
extend_seq_lens_cpu=[],
extend_prefix_lens_cpu=[],
req_pool_indices=[0],
input_ids=torch.tensor([0], dtype=torch.int64),
),
)
for tp_rank in range(self.TP_SIZE):
rank_cca, _ = _make_tiny_cca(
seed=22 + tp_rank, tp_rank=tp_rank, tp_size=self.TP_SIZE
)
self._slice_full_state_dict_into_rank(ref_cca, rank_cca, tp_rank)
rank_pool = _MockReqToTokenPool(
pool_size=8, cca_config=cfg, tp_size=self.TP_SIZE
)
with _mock_pool_context(rank_pool):
rank_cca.forward(
hs_prefill,
_make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S0],
extend_prefix_lens_cpu=[0],
req_pool_indices=[0],
input_ids=torch.arange(S0, dtype=torch.int64),
),
)
rank_q, rank_k, rank_v = rank_cca.forward(
hs_decode,
_make_forward_batch(
is_decode=True,
extend_seq_lens_cpu=[],
extend_prefix_lens_cpu=[],
req_pool_indices=[0],
input_ids=torch.tensor([0], dtype=torch.int64),
),
)
self._check_per_rank_outputs(
full_q, full_k, full_v, rank_q, rank_k, rank_v, tp_rank, cfg
)
def test_tp2_conv_state_is_per_rank_sliced(self):
"""After a TP=2 prefill, each rank's conv state must equal the head
slice of the TP=1 conv state corresponding to that rank's heads.
"""
ref_cca, cfg = _make_tiny_cca(seed=23, tp_rank=0, tp_size=1)
S = 4
torch.manual_seed(903)
hs = torch.randn(S, ref_cca.hidden_size, dtype=torch.float32) * 0.1
ref_pool = _MockReqToTokenPool(pool_size=4, cca_config=cfg, tp_size=1)
with _mock_pool_context(ref_pool):
ref_cca.forward(
hs,
_make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S],
extend_prefix_lens_cpu=[0],
req_pool_indices=[0],
input_ids=torch.arange(S, dtype=torch.int64),
),
)
full_state = ref_pool.mamba2_layer_cache(0).conv[0][0] # [in_out_ch_full, pad]
head_dim = cfg.head_dim
num_q_heads_full = cfg.num_attention_heads
num_k_heads_full = cfg.num_query_groups
latent_q_full = num_q_heads_full * head_dim
q_per_rank = num_q_heads_full // self.TP_SIZE
k_per_rank = num_k_heads_full // self.TP_SIZE
for tp_rank in range(self.TP_SIZE):
rank_cca, _ = _make_tiny_cca(
seed=23 + tp_rank, tp_rank=tp_rank, tp_size=self.TP_SIZE
)
self._slice_full_state_dict_into_rank(ref_cca, rank_cca, tp_rank)
rank_pool = _MockReqToTokenPool(
pool_size=4, cca_config=cfg, tp_size=self.TP_SIZE
)
with _mock_pool_context(rank_pool):
rank_cca.forward(
hs,
_make_forward_batch(
is_decode=False,
extend_seq_lens_cpu=[S],
extend_prefix_lens_cpu=[0],
req_pool_indices=[0],
input_ids=torch.arange(S, dtype=torch.int64),
),
)
rank_state = rank_pool.mamba2_layer_cache(0).conv[0][0]
q_lo = tp_rank * q_per_rank * head_dim
q_hi = q_lo + q_per_rank * head_dim
k_lo = latent_q_full + tp_rank * k_per_rank * head_dim
k_hi = k_lo + k_per_rank * head_dim
expected = torch.cat([full_state[q_lo:q_hi], full_state[k_lo:k_hi]], dim=0)
torch.testing.assert_close(rank_state, expected, atol=1e-5, rtol=1e-5)
def test_tp_assertions_reject_indivisible_head_counts(self):
"""The CCA constructor must reject TP sizes that don't evenly divide
both num_q_heads and num_k_heads, since both grouped-mean and
conv_qk.1 require each rank to hold whole K-head groups.
"""
from sglang.srt.models.zaya import CCA
cfg = _make_tiny_config()
# tiny config has num_query_groups=2; TP=4 cannot divide it cleanly.
with self.assertRaises(AssertionError):
CCA(
config=cfg,
cca_num_k_heads=cfg.num_query_groups,
cca_num_q_heads=cfg.num_attention_heads,
hidden_size=cfg.hidden_size,
head_dim=cfg.head_dim,
cca_time0=cfg.cca_time0,
cca_time1=cfg.cca_time1,
layer_id=0,
tp_rank=0,
tp_size=4,
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,178 @@
"""Numerical correctness test for ZAYA1 MoE + MOD under TP>1.
Background: the MOD (mixture-of-depths) skip-expert residual blend must be
combined with the experts output on the correct side of the cross-rank
all-reduce. ``mod_out = hidden_states * prob`` is replicated on every TP rank,
so all-reducing it would multiply it by ``tp_size``. The model therefore masks
the *per-rank partial* experts output before the reduce and only adds the
replicated ``mod_out`` afterwards:
sum_r(mask · partial_r) + (1 - mask) · mod_out
= mask · experts_out_full + (1 - mask) · mod_out
This test drives the *real* helpers used by ``ZayaBlock.forward`` --
``mod_premask_experts`` and ``mod_blend`` -- so a regression in that math is
caught. The cross-rank all-reduce is simulated by summing the per-rank partials
(the masks are replicated, so the sum is exact), which keeps the test runnable
on CPU CI without a live ``torch.distributed`` group.
"""
import unittest
import torch
from sglang.srt.models.zaya import mod_blend, mod_premask_experts
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
def _reference_blend(
hidden_states: torch.Tensor, # [T, H]
probs: torch.Tensor, # [T, 1]
indices: torch.Tensor, # [T, 1]
experts_out_full: torch.Tensor, # [T, H] -- already-reduced full experts output
num_moe_experts: int,
) -> torch.Tensor:
"""Reference: apply the MOD mask to the *full* (already-reduced) experts
output, then add the skip path. Mirrors the intended algebra directly.
"""
mod_mask = (indices != num_moe_experts).to(experts_out_full.dtype)
mod_out = hidden_states * probs
return mod_mask * experts_out_full + (1.0 - mod_mask) * mod_out
def _real_tp_blend(
hidden_states: torch.Tensor,
probs: torch.Tensor,
indices: torch.Tensor,
partial_experts_per_rank: list[torch.Tensor], # one [T, H] per rank
num_moe_experts: int,
) -> torch.Tensor:
"""Production path: ``mod_premask_experts`` per rank -> simulated all-reduce
(sum) -> ``mod_blend``. Uses the exact helpers ``ZayaBlock.forward`` calls.
"""
mod_out = hidden_states * probs
reduced = None
mod_mask = None
for partial in partial_experts_per_rank:
mask, masked = mod_premask_experts(partial, indices, num_moe_experts)
mod_mask = mask
reduced = masked if reduced is None else reduced + masked
return mod_blend(reduced, mod_mask, mod_out)
def _buggy_old_tp_blend(
hidden_states: torch.Tensor,
probs: torch.Tensor,
indices: torch.Tensor,
partial_experts_per_rank: list[torch.Tensor],
num_moe_experts: int,
) -> torch.Tensor:
"""Old, broken sequence: all-reduce the replicated ``mod_out`` (so it gets
scaled by ``tp_size``) then mix. Proves the test catches a regression.
"""
tp_size = len(partial_experts_per_rank)
mod_out_replicated = hidden_states * probs
mod_out_after_allreduce = mod_out_replicated * tp_size # all-reduce of replicated
experts_out_full = torch.stack(partial_experts_per_rank, dim=0).sum(dim=0)
mod_mask = (indices != num_moe_experts).to(experts_out_full.dtype)
return mod_mask * experts_out_full + (1.0 - mod_mask) * mod_out_after_allreduce
class TestZayaMODUnderTP(CustomTestCase):
def _make_partials(self, T: int, H: int, tp_size: int):
torch.manual_seed(31)
experts_out_full = torch.randn(T, H, dtype=torch.float32) * 0.1
# Split into ``tp_size`` random partial tensors that sum to the full output.
partials = []
remaining = experts_out_full.clone()
for _ in range(tp_size - 1):
p = torch.randn_like(remaining) * 0.05
partials.append(p)
remaining = remaining - p
partials.append(remaining)
return experts_out_full, partials
def _make_inputs(self, T: int, H: int, num_experts: int, frac_skip: float):
torch.manual_seed(7)
hidden_states = torch.randn(T, H, dtype=torch.float32)
probs = torch.rand(T, 1, dtype=torch.float32)
# Build indices: with probability ``frac_skip`` mark token as skip-expert.
skip_id = num_experts # MOD uses ``num_moe_experts`` as the skip slot
rand = torch.rand(T, 1)
real = torch.randint(0, num_experts, (T, 1))
indices = torch.where(rand < frac_skip, torch.full_like(real, skip_id), real)
return hidden_states, probs, indices
def test_real_helpers_match_reference_for_tp(self):
"""The real ``mod_premask_experts`` / ``mod_blend`` path must equal the
reference blend for any TP size and any skip fraction.
"""
T, H = 8, 16
num_experts = 4
for tp_size in (2, 4, 8):
for frac_skip in (0.0, 0.5, 1.0):
hidden_states, probs, indices = self._make_inputs(
T, H, num_experts, frac_skip
)
full, partials = self._make_partials(T, H, tp_size)
ref = _reference_blend(hidden_states, probs, indices, full, num_experts)
real = _real_tp_blend(
hidden_states, probs, indices, partials, num_experts
)
torch.testing.assert_close(
real,
ref,
atol=1e-5,
rtol=1e-5,
msg=f"tp_size={tp_size} frac_skip={frac_skip}",
)
def test_premask_zeroes_skip_tokens(self):
"""``mod_premask_experts`` must zero the experts contribution exactly on
skip-routed tokens and pass through real-expert tokens unchanged.
"""
T, H = 6, 8
num_experts = 4
experts_out = torch.randn(T, H, dtype=torch.float32)
# Alternate skip / real tokens.
indices = torch.tensor(
[[num_experts], [0], [num_experts], [1], [num_experts], [2]],
dtype=torch.long,
)
mod_mask, masked = mod_premask_experts(experts_out, indices, num_experts)
skip_rows = indices.squeeze(-1) == num_experts
self.assertTrue(torch.all(masked[skip_rows] == 0))
torch.testing.assert_close(masked[~skip_rows], experts_out[~skip_rows])
# mask is 0 on skip rows, 1 elsewhere.
self.assertTrue(torch.all(mod_mask.squeeze(-1)[skip_rows] == 0))
self.assertTrue(torch.all(mod_mask.squeeze(-1)[~skip_rows] == 1))
def test_old_blend_is_wrong_when_skip_used(self):
"""Sanity: confirm the old (all-reduce mod_out) formula diverges from the
reference so a regression to that behavior would be caught.
"""
T, H = 8, 16
num_experts = 4
tp_size = 4
hidden_states, probs, indices = self._make_inputs(
T, H, num_experts, frac_skip=0.5
)
full, partials = self._make_partials(T, H, tp_size)
ref = _reference_blend(hidden_states, probs, indices, full, num_experts)
buggy = _buggy_old_tp_blend(
hidden_states, probs, indices, partials, num_experts
)
with self.assertRaises(AssertionError):
torch.testing.assert_close(buggy, ref, atol=1e-3, rtol=1e-3)
if __name__ == "__main__":
unittest.main()