Skip to content

Commit 3b7c649

Browse files
[V1] Warm up seeded and greedy sampler paths (#54425, #54455)
Pre-compile native Gumbel and bfloat16 rejection kernels by cycling stochastic, seeded, and greedy sampling configs during startup warmup. Signed-off-by: Manikanta Bandham <bandhammanikanta@gmail.com>
1 parent 144e79c commit 3b7c649

4 files changed

Lines changed: 59 additions & 4 deletions

File tree

tests/v1/worker/test_gpu_sampler_flags.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -88,3 +88,31 @@ def test_logits_processing_cache_only_checks_active_requests():
8888

8989
assert not np.any(sampler.needs_logits_processing[sampling_only])
9090
assert np.any(sampler.needs_logits_processing[with_processing])
91+
92+
93+
def test_all_sampler_warmup_configs():
94+
"""Test that for_all_sampler_warmup_configs returns a list of configurations
95+
covering:
96+
- FlashInfer (unseeded stochastic),
97+
- native top-k/top-p + Gumbel (seeded), and
98+
- greedy (model dtype) paths."""
99+
configs = SamplingParams.for_all_sampler_warmup_configs()
100+
assert len(configs) >= 3
101+
102+
unseeded = configs[0]
103+
assert unseeded.seed is None
104+
assert unseeded.temperature > 0.0
105+
106+
seeded = configs[1]
107+
assert seeded.seed is not None
108+
109+
greedy = configs[2]
110+
assert greedy.temperature == 0.0
111+
112+
sampler = _make_sampler()
113+
sampler.add_request(0, 1, unseeded)
114+
sampler.add_request(1, 1, seeded)
115+
sampler.add_request(2, 1, greedy)
116+
117+
assert sampler.needs_logits_processing[0]
118+
assert not sampler.needs_logits_processing[2]

vllm/sampling_params.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1270,6 +1270,15 @@ def for_sampler_warmup() -> "SamplingParams":
12701270
prompt_logprobs=1,
12711271
)
12721272

1273+
@classmethod
1274+
def for_all_sampler_warmup_configs(cls) -> list["SamplingParams"]:
1275+
"""Return SamplingParams covering all sampler warmup configurations."""
1276+
return [
1277+
cls.for_sampler_warmup(),
1278+
cls(temperature=0.9, seed=42),
1279+
cls(temperature=0.0),
1280+
]
1281+
12731282

12741283
class BeamSearchParams(
12751284
msgspec.Struct,

vllm/v1/worker/gpu/warmup.py

Lines changed: 14 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -281,14 +281,14 @@ def warmup_kernels(
281281

282282
# SamplingParams exercising all sampling features.
283283
if model_runner.is_pooling_model:
284-
sampling_params = None
284+
sampling_params_list = [None]
285285
pooling_task = model_runner.model_config.get_pooling_task(
286286
model_runner.get_supported_tasks()
287287
)
288288
pooling_params = PoolingParams(task=pooling_task)
289289
pooling_params.verify(model_runner.model_config)
290290
else:
291-
sampling_params = SamplingParams.for_sampler_warmup()
291+
sampling_params_list = SamplingParams.for_all_sampler_warmup_configs()
292292
pooling_params = None
293293

294294
# Assign distinct block IDs per request per group. 0 null block, start from 1.
@@ -309,7 +309,7 @@ def _alloc_blocks(num_blocks: int) -> list[int]:
309309
Request(
310310
req_ids[i],
311311
prompt_token_ids,
312-
sampling_params,
312+
sampling_params_list[i % len(sampling_params_list)],
313313
pooling_params,
314314
mm_features=warmup_mm_features,
315315
),
@@ -410,10 +410,20 @@ def _run_decode_step(indices: list[int], spec_flags: list[bool]) -> None:
410410
# Exercise the model paths that split a batch by whether each
411411
# request received draft tokens.
412412
decode_steps.append(([0, 1], [False, False]))
413-
if num_reqs > 1:
413+
414+
# Single-request steps covering unseeded, seeded, and greedy configs.
414415
decode_steps.append(([0], [use_spec_decode]))
415416
if use_spec_decode:
416417
decode_steps.append(([0], [False]))
418+
419+
decode_steps.append(([1], [use_spec_decode]))
420+
if use_spec_decode:
421+
decode_steps.append(([1], [False]))
422+
423+
if num_reqs >= 3:
424+
decode_steps.append(([2], [use_spec_decode]))
425+
if use_spec_decode:
426+
decode_steps.append(([2], [False]))
417427
elif use_spec_decode:
418428
decode_steps.append(([0], [False]))
419429

vllm/v1/worker/gpu_model_runner.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6446,6 +6446,14 @@ def _dummy_sampler_run(
64466446
logits,
64476447
all_greedy_metadata,
64486448
)
6449+
if self.model_config.dtype != logits.dtype:
6450+
model_dtype_logits = logits.to(self.model_config.dtype)
6451+
self.rejection_sampler(
6452+
dummy_spec_decode_metadata,
6453+
draft_probs,
6454+
model_dtype_logits,
6455+
all_greedy_metadata,
6456+
)
64496457
torch.accelerator.synchronize()
64506458
return sampler_output
64516459

0 commit comments

Comments
 (0)