Files
enginex-ascend-910-vllm/tests/ut/_310p/sample/test_sampler_310.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

284 lines
11 KiB
Python

import sys
import unittest
from contextlib import nullcontext
from pathlib import Path
from types import ModuleType
from unittest.mock import MagicMock, patch
import torch
PROJECT_ROOT = Path(__file__).resolve().parents[4]
if str(PROJECT_ROOT) not in sys.path:
sys.path.insert(0, str(PROJECT_ROOT))
if "vllm" not in sys.modules:
vllm_module = ModuleType("vllm")
vllm_envs_module = ModuleType("vllm.envs")
vllm_envs_module.VLLM_BATCH_INVARIANT = False # type: ignore[attr-defined]
vllm_module.envs = vllm_envs_module # type: ignore[attr-defined]
sys.modules["vllm"] = vllm_module
sys.modules["vllm.envs"] = vllm_envs_module
if "vllm_ascend.sample.sampler" not in sys.modules:
sample_sampler_module = ModuleType("vllm_ascend.sample.sampler")
sample_sampler_module.DEFAULT_LOGPROBS_MODE = "raw_logprobs" # type: ignore[attr-defined]
sample_sampler_module.AscendSampler = type("AscendSampler", (), {}) # type: ignore[attr-defined]
sample_sampler_module.AscendTopKTopPSampler = type("AscendTopKTopPSampler", (), {}) # type: ignore[attr-defined]
sys.modules["vllm_ascend.sample.sampler"] = sample_sampler_module
if "vllm_ascend.utils" not in sys.modules:
utils_module = ModuleType("vllm_ascend.utils")
utils_module.global_stream = lambda: MagicMock() # type: ignore[attr-defined]
utils_module.npu_stream_switch = lambda _: nullcontext() # type: ignore[attr-defined]
sys.modules["vllm_ascend.utils"] = utils_module
from vllm_ascend._310p.sample import sampler as sampler_310p # noqa: E402
class _FakeRow:
def __init__(self):
self.generators = []
def exponential_(self, generator=None):
self.generators.append(generator)
return self
class _FakeQ:
def __init__(self, batch_size):
self.shape = (batch_size, 4)
self.default_exponential_called = False
self.rows = {idx: _FakeRow() for idx in range(batch_size)}
def cpu(self):
return self
def npu(self):
return self
def exponential_(self, generator=None):
if generator is None:
self.default_exponential_called = True
return self
def __getitem__(self, idx):
return self.rows[idx]
def __setitem__(self, idx, value):
self.rows[idx] = value
def _empty_like_side_effect(q_instances, template):
if isinstance(template, _FakeRow):
return _FakeRow()
return next(q_instances)
class _FakeCPUGenerator:
def __init__(self, device=None):
self.device = device
self.state = None
self.seed = None
def set_state(self, state):
self.state = state
def manual_seed(self, seed):
self.seed = seed
class TestSampler310pStandalone(unittest.TestCase):
def tearDown(self):
sampler_310p._CPU_GENERATOR_CACHE_310P.clear()
def test_random_sample_310p_reuse_cpu_generator_cache(self):
sampler_310p._CPU_GENERATOR_CACHE_310P.clear()
probs = MagicMock()
probs.div_.return_value = probs
probs.argmax.return_value = probs
probs.view.return_value = torch.tensor([0])
fake_q_first = _FakeQ(batch_size=2)
fake_q_second = _FakeQ(batch_size=2)
q_instances = iter([fake_q_first, fake_q_second])
npu_stream = MagicMock()
generator = MagicMock()
generator.get_state.return_value = b"state"
generator.initial_seed.return_value = 7
generators = {1: generator}
with (
patch.object(sampler_310p, "npu_stream_switch", return_value=nullcontext()),
patch.object(sampler_310p, "global_stream", return_value=MagicMock()),
patch.object(
sampler_310p.torch,
"empty_like",
side_effect=lambda template: _empty_like_side_effect(q_instances, template),
),
patch.object(sampler_310p.torch, "Generator", side_effect=_FakeCPUGenerator) as gen_ctor,
patch.object(
sampler_310p.torch,
"npu",
ModuleType("torch.npu"),
create=True,
),
):
sampler_310p.torch.npu.current_stream = MagicMock(return_value=npu_stream)
sampler_310p._random_sample_310p(probs, generators)
sampler_310p._random_sample_310p(probs, generators)
self.assertEqual(gen_ctor.call_count, 1)
self.assertIn(1, sampler_310p._CPU_GENERATOR_CACHE_310P)
cached_cpu_generator, source_generator_id = sampler_310p._CPU_GENERATOR_CACHE_310P[1]
self.assertIs(fake_q_first.rows[1].generators[0], cached_cpu_generator)
self.assertIs(fake_q_second.rows[1].generators[0], cached_cpu_generator)
self.assertEqual(source_generator_id, id(generator))
self.assertEqual(cached_cpu_generator.state, b"state")
self.assertIsNone(cached_cpu_generator.seed)
self.assertEqual(npu_stream.wait_stream.call_count, 2)
def test_random_sample_310p_fallback_to_initial_seed_when_set_state_failed(self):
sampler_310p._CPU_GENERATOR_CACHE_310P.clear()
probs = MagicMock()
probs.div_.return_value = probs
probs.argmax.return_value = probs
probs.view.return_value = torch.tensor([1])
fake_q = _FakeQ(batch_size=1)
q_instances = iter([fake_q])
npu_stream = MagicMock()
generator = MagicMock()
generator.get_state.side_effect = RuntimeError("state read failed")
generator.initial_seed.return_value = 1234
generators = {0: generator}
class _FailSetStateCPUGenerator(_FakeCPUGenerator):
def set_state(self, state):
raise RuntimeError("state set failed")
with (
patch.object(sampler_310p, "npu_stream_switch", return_value=nullcontext()),
patch.object(sampler_310p, "global_stream", return_value=MagicMock()),
patch.object(
sampler_310p.torch,
"empty_like",
side_effect=lambda template: _empty_like_side_effect(q_instances, template),
),
patch.object(sampler_310p.torch, "Generator", side_effect=_FailSetStateCPUGenerator),
patch.object(
sampler_310p.torch,
"npu",
ModuleType("torch.npu"),
create=True,
),
):
sampler_310p.torch.npu.current_stream = MagicMock(return_value=npu_stream)
sampler_310p._random_sample_310p(probs, generators)
cached_cpu_generator, source_generator_id = sampler_310p._CPU_GENERATOR_CACHE_310P[0]
self.assertEqual(source_generator_id, id(generator))
self.assertEqual(cached_cpu_generator.seed, 1234)
self.assertIs(fake_q.rows[0].generators[0], cached_cpu_generator)
self.assertEqual(npu_stream.wait_stream.call_count, 1)
def test_random_sample_310p_rebuild_cache_when_generator_identity_changes(self):
sampler_310p._CPU_GENERATOR_CACHE_310P.clear()
probs = MagicMock()
probs.div_.return_value = probs
probs.argmax.return_value = probs
probs.view.return_value = torch.tensor([0])
fake_q_first = _FakeQ(batch_size=1)
fake_q_second = _FakeQ(batch_size=1)
q_instances = iter([fake_q_first, fake_q_second])
npu_stream = MagicMock()
generator_first = MagicMock()
generator_first.get_state.return_value = b"state-1"
generator_first.initial_seed.return_value = 11
generator_second = MagicMock()
generator_second.get_state.return_value = b"state-2"
generator_second.initial_seed.return_value = 22
with (
patch.object(sampler_310p, "npu_stream_switch", return_value=nullcontext()),
patch.object(sampler_310p, "global_stream", return_value=MagicMock()),
patch.object(
sampler_310p.torch,
"empty_like",
side_effect=lambda template: _empty_like_side_effect(q_instances, template),
),
patch.object(sampler_310p.torch, "Generator", side_effect=_FakeCPUGenerator) as gen_ctor,
patch.object(
sampler_310p.torch,
"npu",
ModuleType("torch.npu"),
create=True,
),
):
sampler_310p.torch.npu.current_stream = MagicMock(return_value=npu_stream)
sampler_310p._random_sample_310p(probs, {0: generator_first})
sampler_310p._random_sample_310p(probs, {0: generator_second})
self.assertEqual(gen_ctor.call_count, 2)
first_cpu_generator = fake_q_first.rows[0].generators[0]
second_cpu_generator = fake_q_second.rows[0].generators[0]
self.assertIsNot(first_cpu_generator, second_cpu_generator)
self.assertEqual(first_cpu_generator.state, b"state-1")
self.assertEqual(second_cpu_generator.state, b"state-2")
cached_cpu_generator, source_generator_id = sampler_310p._CPU_GENERATOR_CACHE_310P[0]
self.assertIs(cached_cpu_generator, second_cpu_generator)
self.assertEqual(source_generator_id, id(generator_second))
def test_fill_cpu_exponential_310p_moves_has_draft_mask_to_cpu(self):
"""Regression: NPU has_draft_mask must be moved to CPU before torch.where."""
sampler_310p._CPU_GENERATOR_CACHE_310P.clear()
q_cpu = torch.full((2, 4), 7.0)
cpu_mask = torch.tensor([True, False])
has_draft_mask = MagicMock()
has_draft_mask.cpu.return_value = cpu_mask
def _make_source_generator(seed: int):
source_generator = MagicMock()
seed_generator = torch.Generator(device="cpu")
seed_generator.manual_seed(seed)
source_generator.get_state.return_value = seed_generator.get_state()
source_generator.initial_seed.return_value = seed
return source_generator
where_conditions = []
real_where = torch.where
def where_spy(condition, x, y):
where_conditions.append(condition.detach().clone())
self.assertEqual(condition.device.type, "cpu")
self.assertEqual(x.device.type, "cpu")
self.assertEqual(y.device.type, "cpu")
return real_where(condition, x, y)
with patch.object(sampler_310p.torch, "where", side_effect=where_spy):
sampler_310p._fill_cpu_exponential_310p(
q_cpu,
{
0: _make_source_generator(42),
1: _make_source_generator(43),
},
has_draft_mask,
)
has_draft_mask.cpu.assert_called_once()
self.assertEqual(len(where_conditions), 2)
self.assertTrue(bool(where_conditions[0]))
self.assertFalse(bool(where_conditions[1]))
# Row 0 (masked): overwritten by seeded exponential via torch.where.
self.assertFalse(torch.equal(q_cpu[0], torch.full((4,), 7.0)))
# Row 1 (unmasked): also overwritten by the default exponential_ prefill.
self.assertFalse(torch.equal(q_cpu[1], torch.full((4,), 7.0)))
if __name__ == "__main__":
unittest.main()