Skip to content

Commit 32a59d4

Browse files
authored
[Ops][Misc] Optimize 310P random sampling with inverse CDF (vllm-project#12966)
### What this PR does / why we need it? This PR optimizes temperature-based random sampling on 310P without generating random values on the NPU. For regular sampling, it replaces the `[batch, vocab]` CPU exponential-noise path and `Div + ArgMax` with inverse-CDF sampling: one CPU uniform random value per request, asynchronous H2D transfer, followed by NPU `Cumsum + SearchSorted`. For MTP recovered-token sampling, exponential values are still generated on CPU via `-log(U)`, with improved request-specific generator reuse, inactive-request handling, and removal of redundant copies and host-side stream synchronization. This avoids the problematic 310P NPU random/in-place Add path while significantly reducing CPU RNG and H2D overhead. ### Does this PR introduce _any_ user-facing change? NA ### How was this patch tested? Local test RC without temperature (benchmark): <img width="1205" height="784" alt="image" src="https://github.com/user-attachments/assets/62269e4d-4275-4b78-8249-45ecffa1b3b9" /> RC with temperature before this PR: <img width="1193" height="778" alt="image1" src="https://github.com/user-attachments/assets/f4d75b4b-5994-488e-8a0d-e1db1cbe80f8" /> RC with temperature after this PR: <img width="1183" height="784" alt="image2" src="https://github.com/user-attachments/assets/e90da9a1-695b-4201-9747-68724053b517" /> - vLLM version: v0.25.1 - vLLM main: vllm-project/vllm@d02df74 --------- Signed-off-by: shysummer <sunhaoyu14@huawei.com> Signed-off-by: sunhaoyu <sunhaoyu14@huawei.com>
1 parent b6fe31d commit 32a59d4

3 files changed

Lines changed: 244 additions & 270 deletions

File tree

tests/ut/_310p/sample/test_sampler_310.py

Lines changed: 131 additions & 218 deletions
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,9 @@
2323
sample_sampler_module = ModuleType("vllm_ascend.sample.sampler")
2424
sample_sampler_module.DEFAULT_LOGPROBS_MODE = "raw_logprobs" # type: ignore[attr-defined]
2525
sample_sampler_module.AscendSampler = type("AscendSampler", (), {}) # type: ignore[attr-defined]
26-
sample_sampler_module.AscendTopKTopPSampler = type("AscendTopKTopPSampler", (), {}) # type: ignore[attr-defined]
26+
sample_sampler_module.AscendTopKTopPSampler = type( # type: ignore[attr-defined]
27+
"AscendTopKTopPSampler", (), {}
28+
)
2729
sys.modules["vllm_ascend.sample.sampler"] = sample_sampler_module
2830

2931
if "vllm_ascend.utils" not in sys.modules:
@@ -35,248 +37,159 @@
3537
from vllm_ascend._310p.sample import sampler as sampler_310p # noqa: E402
3638

3739

38-
class _FakeRow:
39-
def __init__(self):
40-
self.generators = []
41-
42-
def exponential_(self, generator=None):
43-
self.generators.append(generator)
44-
return self
45-
46-
47-
class _FakeQ:
48-
def __init__(self, batch_size):
49-
self.shape = (batch_size, 4)
50-
self.default_exponential_called = False
51-
self.rows = {idx: _FakeRow() for idx in range(batch_size)}
52-
53-
def cpu(self):
54-
return self
55-
56-
def npu(self):
57-
return self
58-
59-
def exponential_(self, generator=None):
60-
if generator is None:
61-
self.default_exponential_called = True
62-
return self
63-
64-
def __getitem__(self, idx):
65-
return self.rows[idx]
66-
67-
def __setitem__(self, idx, value):
68-
self.rows[idx] = value
40+
class _SourceGenerator:
41+
def __init__(self, seed: int):
42+
self.seed = seed
6943

44+
def initial_seed(self) -> int:
45+
return self.seed
7046

71-
def _empty_like_side_effect(q_instances, template):
72-
if isinstance(template, _FakeRow):
73-
return _FakeRow()
74-
return next(q_instances)
7547

48+
class _RecordingStreamContext:
49+
def __init__(self, events: list[str]):
50+
self.events = events
7651

77-
class _FakeCPUGenerator:
78-
def __init__(self, device=None):
79-
self.device = device
80-
self.state = None
81-
self.seed = None
52+
def __enter__(self):
53+
self.events.append("enter_global")
8254

83-
def set_state(self, state):
84-
self.state = state
85-
86-
def manual_seed(self, seed):
87-
self.seed = seed
55+
def __exit__(self, exc_type, exc_value, traceback):
56+
self.events.append("exit_global")
8857

8958

9059
class TestSampler310pStandalone(unittest.TestCase):
9160
def tearDown(self):
9261
sampler_310p._CPU_GENERATOR_CACHE_310P.clear()
9362

94-
def test_random_sample_310p_reuse_cpu_generator_cache(self):
95-
sampler_310p._CPU_GENERATOR_CACHE_310P.clear()
63+
def test_prepare_cpu_generators_preserves_requests_across_reordering(self):
64+
source_a = _SourceGenerator(11)
65+
source_b = _SourceGenerator(22)
66+
67+
first = sampler_310p._prepare_cpu_generators_310p({0: source_a, 1: source_b})
68+
second = sampler_310p._prepare_cpu_generators_310p({0: source_b, 1: source_a})
69+
70+
self.assertIs(second[1], first[0])
71+
self.assertIs(second[0], first[1])
72+
self.assertIs(sampler_310p._CPU_GENERATOR_CACHE_310P[0][0], source_b)
73+
self.assertIs(sampler_310p._CPU_GENERATOR_CACHE_310P[1][0], source_a)
74+
75+
def test_prepare_cpu_generators_replaces_changed_source(self):
76+
source_first = _SourceGenerator(11)
77+
source_second = _SourceGenerator(22)
78+
79+
first = sampler_310p._prepare_cpu_generators_310p({0: source_first})
80+
second = sampler_310p._prepare_cpu_generators_310p({0: source_second})
81+
82+
self.assertIsNot(second[0], first[0])
83+
expected = torch.Generator(device="cpu")
84+
expected.manual_seed(22)
85+
self.assertEqual(
86+
torch.rand((), generator=second[0]).item(),
87+
torch.rand((), generator=expected).item(),
88+
)
89+
90+
def test_sample_from_cdf_handles_zero_weight_prefix_and_boundaries(self):
91+
weights = torch.tensor(
92+
[
93+
[0.0, 2.0, 3.0, 5.0],
94+
[1.0, 0.0, 0.0, 0.0],
95+
[0.0, 0.0, 0.0, 4.0],
96+
]
97+
)
98+
uniforms = torch.tensor([0.2, 0.999, torch.finfo(torch.float32).tiny])
99+
100+
sampled = sampler_310p._sample_from_cdf_310p(weights, uniforms)
101+
102+
self.assertTrue(torch.equal(sampled, torch.tensor([2, 0, 3])))
103+
104+
def test_fill_exponential_honors_active_mask(self):
105+
source_active = _SourceGenerator(31)
106+
source_inactive = _SourceGenerator(41)
107+
prepared = sampler_310p._prepare_cpu_generators_310p({0: source_active, 1: source_inactive})
108+
active_state = prepared[0].get_state()
109+
inactive_state = prepared[1].get_state()
110+
real_rand = torch.rand
111+
112+
def rand_without_pinning(*args, **kwargs):
113+
kwargs.pop("pin_memory", None)
114+
return real_rand(*args, **kwargs)
115+
116+
with patch.object(sampler_310p.torch, "rand", side_effect=rand_without_pinning):
117+
exponential = sampler_310p.fill_exponential_310p(
118+
torch.empty((2, 4), dtype=torch.float32),
119+
{0: source_active, 1: source_inactive},
120+
active_mask=[True, False],
121+
)
122+
123+
cached_active = sampler_310p._CPU_GENERATOR_CACHE_310P[0][1]
124+
cached_inactive = sampler_310p._CPU_GENERATOR_CACHE_310P[1][1]
125+
self.assertFalse(torch.equal(cached_active.get_state(), active_state))
126+
self.assertTrue(torch.equal(cached_inactive.get_state(), inactive_state))
127+
self.assertEqual(exponential.shape, (2, 4))
128+
self.assertTrue(torch.isfinite(exponential).all())
129+
self.assertTrue((exponential > 0).all())
130+
131+
def test_random_sample_waits_before_cdf_reads_probs(self):
132+
events: list[str] = []
133+
global_npu_stream = MagicMock()
134+
current_npu_stream = MagicMock()
135+
current_npu_stream.wait_stream.side_effect = lambda _: events.append("wait_global")
136+
fake_npu = ModuleType("torch.npu")
137+
fake_npu.current_stream = MagicMock( # type: ignore[attr-defined]
138+
return_value=current_npu_stream
139+
)
96140
probs = MagicMock()
97-
probs.div_.return_value = probs
98-
probs.argmax.return_value = probs
99-
probs.view.return_value = torch.tensor([0])
141+
probs.shape = (1, 4)
142+
probs.device = torch.device("cpu")
143+
uniforms = MagicMock()
100144

101-
fake_q_first = _FakeQ(batch_size=2)
102-
fake_q_second = _FakeQ(batch_size=2)
103-
q_instances = iter([fake_q_first, fake_q_second])
145+
def generate_uniforms(*args, **kwargs):
146+
events.append("generate_uniforms")
147+
return uniforms
104148

105-
npu_stream = MagicMock()
106-
generator = MagicMock()
107-
generator.get_state.return_value = b"state"
108-
generator.initial_seed.return_value = 7
109-
generators = {1: generator}
149+
def sample_from_cdf(actual_probs, actual_uniforms):
150+
events.append("sample_from_cdf")
151+
self.assertIs(actual_probs, probs)
152+
self.assertIs(actual_uniforms, uniforms)
153+
return torch.tensor([2])
110154

111155
with (
112-
patch.object(sampler_310p, "npu_stream_switch", return_value=nullcontext()),
113-
patch.object(sampler_310p, "global_stream", return_value=MagicMock()),
114156
patch.object(
115-
sampler_310p.torch,
116-
"empty_like",
117-
side_effect=lambda template: _empty_like_side_effect(q_instances, template),
157+
sampler_310p,
158+
"npu_stream_switch",
159+
return_value=_RecordingStreamContext(events),
118160
),
119-
patch.object(sampler_310p.torch, "Generator", side_effect=_FakeCPUGenerator) as gen_ctor,
120161
patch.object(
121-
sampler_310p.torch,
122-
"npu",
123-
ModuleType("torch.npu"),
124-
create=True,
162+
sampler_310p,
163+
"global_stream",
164+
return_value=global_npu_stream,
125165
),
126-
):
127-
sampler_310p.torch.npu.current_stream = MagicMock(return_value=npu_stream)
128-
sampler_310p._random_sample_310p(probs, generators)
129-
sampler_310p._random_sample_310p(probs, generators)
130-
131-
self.assertEqual(gen_ctor.call_count, 1)
132-
self.assertIn(1, sampler_310p._CPU_GENERATOR_CACHE_310P)
133-
cached_cpu_generator, source_generator_id = sampler_310p._CPU_GENERATOR_CACHE_310P[1]
134-
self.assertIs(fake_q_first.rows[1].generators[0], cached_cpu_generator)
135-
self.assertIs(fake_q_second.rows[1].generators[0], cached_cpu_generator)
136-
self.assertEqual(source_generator_id, id(generator))
137-
self.assertEqual(cached_cpu_generator.state, b"state")
138-
self.assertIsNone(cached_cpu_generator.seed)
139-
self.assertEqual(npu_stream.wait_stream.call_count, 2)
140-
141-
def test_random_sample_310p_fallback_to_initial_seed_when_set_state_failed(self):
142-
sampler_310p._CPU_GENERATOR_CACHE_310P.clear()
143-
probs = MagicMock()
144-
probs.div_.return_value = probs
145-
probs.argmax.return_value = probs
146-
probs.view.return_value = torch.tensor([1])
147-
148-
fake_q = _FakeQ(batch_size=1)
149-
q_instances = iter([fake_q])
150-
npu_stream = MagicMock()
151-
generator = MagicMock()
152-
generator.get_state.side_effect = RuntimeError("state read failed")
153-
generator.initial_seed.return_value = 1234
154-
generators = {0: generator}
155-
156-
class _FailSetStateCPUGenerator(_FakeCPUGenerator):
157-
def set_state(self, state):
158-
raise RuntimeError("state set failed")
159-
160-
with (
161-
patch.object(sampler_310p, "npu_stream_switch", return_value=nullcontext()),
162-
patch.object(sampler_310p, "global_stream", return_value=MagicMock()),
163166
patch.object(
164-
sampler_310p.torch,
165-
"empty_like",
166-
side_effect=lambda template: _empty_like_side_effect(q_instances, template),
167+
sampler_310p,
168+
"_generate_request_uniforms_310p",
169+
side_effect=generate_uniforms,
167170
),
168-
patch.object(sampler_310p.torch, "Generator", side_effect=_FailSetStateCPUGenerator),
169171
patch.object(
170-
sampler_310p.torch,
171-
"npu",
172-
ModuleType("torch.npu"),
173-
create=True,
172+
sampler_310p,
173+
"_sample_from_cdf_310p",
174+
side_effect=sample_from_cdf,
174175
),
176+
patch.object(sampler_310p.torch, "npu", fake_npu, create=True),
175177
):
176-
sampler_310p.torch.npu.current_stream = MagicMock(return_value=npu_stream)
177-
sampler_310p._random_sample_310p(probs, generators)
178-
179-
cached_cpu_generator, source_generator_id = sampler_310p._CPU_GENERATOR_CACHE_310P[0]
180-
self.assertEqual(source_generator_id, id(generator))
181-
self.assertEqual(cached_cpu_generator.seed, 1234)
182-
self.assertIs(fake_q.rows[0].generators[0], cached_cpu_generator)
183-
self.assertEqual(npu_stream.wait_stream.call_count, 1)
184-
185-
def test_random_sample_310p_rebuild_cache_when_generator_identity_changes(self):
186-
sampler_310p._CPU_GENERATOR_CACHE_310P.clear()
187-
probs = MagicMock()
188-
probs.div_.return_value = probs
189-
probs.argmax.return_value = probs
190-
probs.view.return_value = torch.tensor([0])
191-
192-
fake_q_first = _FakeQ(batch_size=1)
193-
fake_q_second = _FakeQ(batch_size=1)
194-
q_instances = iter([fake_q_first, fake_q_second])
195-
npu_stream = MagicMock()
196-
197-
generator_first = MagicMock()
198-
generator_first.get_state.return_value = b"state-1"
199-
generator_first.initial_seed.return_value = 11
200-
201-
generator_second = MagicMock()
202-
generator_second.get_state.return_value = b"state-2"
203-
generator_second.initial_seed.return_value = 22
204-
205-
with (
206-
patch.object(sampler_310p, "npu_stream_switch", return_value=nullcontext()),
207-
patch.object(sampler_310p, "global_stream", return_value=MagicMock()),
208-
patch.object(
209-
sampler_310p.torch,
210-
"empty_like",
211-
side_effect=lambda template: _empty_like_side_effect(q_instances, template),
212-
),
213-
patch.object(sampler_310p.torch, "Generator", side_effect=_FakeCPUGenerator) as gen_ctor,
214-
patch.object(
215-
sampler_310p.torch,
216-
"npu",
217-
ModuleType("torch.npu"),
218-
create=True,
219-
),
220-
):
221-
sampler_310p.torch.npu.current_stream = MagicMock(return_value=npu_stream)
222-
sampler_310p._random_sample_310p(probs, {0: generator_first})
223-
sampler_310p._random_sample_310p(probs, {0: generator_second})
224-
225-
self.assertEqual(gen_ctor.call_count, 2)
226-
first_cpu_generator = fake_q_first.rows[0].generators[0]
227-
second_cpu_generator = fake_q_second.rows[0].generators[0]
228-
self.assertIsNot(first_cpu_generator, second_cpu_generator)
229-
self.assertEqual(first_cpu_generator.state, b"state-1")
230-
self.assertEqual(second_cpu_generator.state, b"state-2")
231-
cached_cpu_generator, source_generator_id = sampler_310p._CPU_GENERATOR_CACHE_310P[0]
232-
self.assertIs(cached_cpu_generator, second_cpu_generator)
233-
self.assertEqual(source_generator_id, id(generator_second))
234-
235-
def test_fill_cpu_exponential_310p_moves_has_draft_mask_to_cpu(self):
236-
"""Regression: NPU has_draft_mask must be moved to CPU before torch.where."""
237-
sampler_310p._CPU_GENERATOR_CACHE_310P.clear()
238-
239-
q_cpu = torch.full((2, 4), 7.0)
240-
cpu_mask = torch.tensor([True, False])
241-
has_draft_mask = MagicMock()
242-
has_draft_mask.cpu.return_value = cpu_mask
243-
244-
def _make_source_generator(seed: int):
245-
source_generator = MagicMock()
246-
seed_generator = torch.Generator(device="cpu")
247-
seed_generator.manual_seed(seed)
248-
source_generator.get_state.return_value = seed_generator.get_state()
249-
source_generator.initial_seed.return_value = seed
250-
return source_generator
251-
252-
where_conditions = []
253-
real_where = torch.where
254-
255-
def where_spy(condition, x, y):
256-
where_conditions.append(condition.detach().clone())
257-
self.assertEqual(condition.device.type, "cpu")
258-
self.assertEqual(x.device.type, "cpu")
259-
self.assertEqual(y.device.type, "cpu")
260-
return real_where(condition, x, y)
261-
262-
with patch.object(sampler_310p.torch, "where", side_effect=where_spy):
263-
sampler_310p._fill_cpu_exponential_310p(
264-
q_cpu,
265-
{
266-
0: _make_source_generator(42),
267-
1: _make_source_generator(43),
268-
},
269-
has_draft_mask,
270-
)
271-
272-
has_draft_mask.cpu.assert_called_once()
273-
self.assertEqual(len(where_conditions), 2)
274-
self.assertTrue(bool(where_conditions[0]))
275-
self.assertFalse(bool(where_conditions[1]))
276-
# Row 0 (masked): overwritten by seeded exponential via torch.where.
277-
self.assertFalse(torch.equal(q_cpu[0], torch.full((4,), 7.0)))
278-
# Row 1 (unmasked): also overwritten by the default exponential_ prefill.
279-
self.assertFalse(torch.equal(q_cpu[1], torch.full((4,), 7.0)))
178+
sampled = sampler_310p._random_sample_310p(probs, {})
179+
180+
self.assertTrue(torch.equal(sampled, torch.tensor([2])))
181+
self.assertEqual(
182+
events,
183+
[
184+
"enter_global",
185+
"generate_uniforms",
186+
"exit_global",
187+
"wait_global",
188+
"sample_from_cdf",
189+
],
190+
)
191+
current_npu_stream.wait_stream.assert_called_once_with(global_npu_stream)
192+
global_npu_stream.wait_stream.assert_not_called()
280193

281194

282195
if __name__ == "__main__":

0 commit comments

Comments
 (0)