Files
sglang/test/registered/unit/managers/test_bailing_modality_metadata.py
T

198 lines
6.7 KiB
Python

"""Regression tests for Bailing modality metadata and per-token routing bias."""
import unittest
from types import SimpleNamespace
import torch
from sglang.srt.configs.model_config import requires_mm_token_modalities
from sglang.srt.layers.moe.topk import biased_grouped_topk_impl
from sglang.srt.layers.multi_gate import create_multi_gate_mm_indices
from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalInputs,
MultimodalProcessorOutput,
)
from sglang.srt.model_executor.forward_batch_info import (
_build_forward_token_modalities,
_maybe_build_forward_token_modalities,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
class TestBailingModalityMetadata(CustomTestCase):
def test_offsets_survive_hash_padding(self):
"""Hash-derived token replacement must not erase image/audio identity."""
items = [
MultimodalDataItem(
modality=Modality.IMAGE,
offsets=[(1, 2)],
feature=torch.ones(1),
),
MultimodalDataItem(
modality=Modality.AUDIO,
offsets=[(4, 5)],
feature=torch.ones(1),
),
]
output = MultimodalProcessorOutput(
mm_items=items,
input_ids=[100, 11, 11, 101, 12, 12, 102],
)
inputs = MultimodalInputs.from_processor_output(
output, requires_mm_token_modalities=True
)
self.assertEqual(
inputs.token_modalities,
[
0,
Modality.IMAGE.value,
Modality.IMAGE.value,
0,
Modality.AUDIO.value,
Modality.AUDIO.value,
0,
],
)
self.assertNotEqual(items[0].pad_value, 11)
self.assertNotEqual(items[1].pad_value, 12)
def test_offset_validation_only_runs_for_multirouter(self):
item = MultimodalDataItem(
modality=Modality.IMAGE,
offsets=[(1, 3)],
feature=torch.ones(1),
)
output = MultimodalProcessorOutput(mm_items=[item], input_ids=[100, 101])
inputs = MultimodalInputs.from_processor_output(output)
self.assertIsNone(inputs.token_modalities)
with self.assertRaisesRegex(ValueError, "Invalid multimodal token offsets"):
MultimodalInputs.from_processor_output(
output, requires_mm_token_modalities=True
)
def test_chunked_metadata_is_identical_on_every_pp_stage(self):
"""Each PP stage must independently receive the same active token map."""
mm_inputs = [
MultimodalInputs(
mm_items=[],
token_modalities=[0, Modality.IMAGE.value, Modality.IMAGE.value, 0],
),
MultimodalInputs(
mm_items=[],
token_modalities=[Modality.AUDIO.value, Modality.AUDIO.value, 0],
),
]
expected = torch.tensor(
[
Modality.IMAGE.value,
Modality.IMAGE.value,
0,
Modality.AUDIO.value,
Modality.AUDIO.value,
],
dtype=torch.int8,
)
stage_maps = [
_build_forward_token_modalities(
mm_inputs,
extend_prefix_lens=[1, 0],
extend_seq_lens=[3, 2],
num_tokens=5,
device=torch.device("cpu"),
)
for _ in range(2)
]
for stage_map in stage_maps:
torch.testing.assert_close(stage_map, expected)
def test_only_bailing_multirouter_requires_token_modalities(self):
bailing_arch = ["BailingMoeV3VLForConditionalGeneration"]
self.assertFalse(
requires_mm_token_modalities(
bailing_arch, SimpleNamespace(multi_gate=False, router_type="topN")
)
)
self.assertTrue(
requires_mm_token_modalities(
bailing_arch, SimpleNamespace(multi_gate=True, router_type="topN")
)
)
self.assertFalse(
requires_mm_token_modalities(
["DeepseekV4ForCausalLM"],
SimpleNamespace(multi_gate=True, router_type="MultiRouter"),
)
)
def test_unrelated_model_skips_mismatched_metadata(self):
mm_inputs = [
MultimodalInputs(mm_items=[], token_modalities=[Modality.IMAGE.value])
]
result = _maybe_build_forward_token_modalities(
SimpleNamespace(requires_mm_token_modalities=False),
mm_inputs,
extend_prefix_lens=[0],
extend_seq_lens=[1],
num_tokens=6,
device=torch.device("cpu"),
)
self.assertIsNone(result)
with self.assertRaisesRegex(ValueError, "does not match the forward batch"):
_maybe_build_forward_token_modalities(
SimpleNamespace(requires_mm_token_modalities=True),
mm_inputs,
extend_prefix_lens=[0],
extend_seq_lens=[1],
num_tokens=6,
device=torch.device("cpu"),
)
def test_mixed_modalities_select_reference_experts(self):
"""Per-token bias must select image/audio experts after modality grouping."""
modalities = torch.tensor(
[
Modality.IMAGE.value,
Modality.IMAGE.value,
0,
Modality.AUDIO.value,
Modality.AUDIO.value,
],
dtype=torch.int8,
)
token_indices, modality_ids = create_multi_gate_mm_indices(modalities)
self.assertEqual(modality_ids.tolist(), [0, 1, 2])
self.assertEqual(token_indices[:1].tolist(), [2])
self.assertEqual(token_indices[64:66].tolist(), [0, 1])
self.assertEqual(token_indices[128:130].tolist(), [3, 4])
router_logits = torch.zeros(5, 8)
dynamic_bias = torch.zeros_like(router_logits)
expected_experts = torch.tensor([1, 1, 0, 6, 6], dtype=torch.int32)
dynamic_bias.scatter_(1, expected_experts.long().unsqueeze(1), 10.0)
_, expert_ids = biased_grouped_topk_impl(
hidden_states=torch.zeros(5, 4),
gating_output=router_logits,
correction_bias=dynamic_bias,
topk=1,
renormalize=True,
num_expert_group=2,
topk_group=1,
)
torch.testing.assert_close(expert_ids.squeeze(1), expected_experts)
if __name__ == "__main__":
unittest.main()