|
23 | 23 | sample_sampler_module = ModuleType("vllm_ascend.sample.sampler") |
24 | 24 | sample_sampler_module.DEFAULT_LOGPROBS_MODE = "raw_logprobs" # type: ignore[attr-defined] |
25 | 25 | 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 | + ) |
27 | 29 | sys.modules["vllm_ascend.sample.sampler"] = sample_sampler_module |
28 | 30 |
|
29 | 31 | if "vllm_ascend.utils" not in sys.modules: |
|
35 | 37 | from vllm_ascend._310p.sample import sampler as sampler_310p # noqa: E402 |
36 | 38 |
|
37 | 39 |
|
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 |
69 | 43 |
|
| 44 | + def initial_seed(self) -> int: |
| 45 | + return self.seed |
70 | 46 |
|
71 | | -def _empty_like_side_effect(q_instances, template): |
72 | | - if isinstance(template, _FakeRow): |
73 | | - return _FakeRow() |
74 | | - return next(q_instances) |
75 | 47 |
|
| 48 | +class _RecordingStreamContext: |
| 49 | + def __init__(self, events: list[str]): |
| 50 | + self.events = events |
76 | 51 |
|
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") |
82 | 54 |
|
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") |
88 | 57 |
|
89 | 58 |
|
90 | 59 | class TestSampler310pStandalone(unittest.TestCase): |
91 | 60 | def tearDown(self): |
92 | 61 | sampler_310p._CPU_GENERATOR_CACHE_310P.clear() |
93 | 62 |
|
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 | + ) |
96 | 140 | 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() |
100 | 144 |
|
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 |
104 | 148 |
|
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]) |
110 | 154 |
|
111 | 155 | with ( |
112 | | - patch.object(sampler_310p, "npu_stream_switch", return_value=nullcontext()), |
113 | | - patch.object(sampler_310p, "global_stream", return_value=MagicMock()), |
114 | 156 | 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), |
118 | 160 | ), |
119 | | - patch.object(sampler_310p.torch, "Generator", side_effect=_FakeCPUGenerator) as gen_ctor, |
120 | 161 | 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, |
125 | 165 | ), |
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()), |
163 | 166 | 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, |
167 | 170 | ), |
168 | | - patch.object(sampler_310p.torch, "Generator", side_effect=_FailSetStateCPUGenerator), |
169 | 171 | 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, |
174 | 175 | ), |
| 176 | + patch.object(sampler_310p.torch, "npu", fake_npu, create=True), |
175 | 177 | ): |
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() |
280 | 193 |
|
281 | 194 |
|
282 | 195 | if __name__ == "__main__": |
|
0 commit comments