From 5ea76fa89aa2e307f0d9a2e7fc19d13aed65a82f Mon Sep 17 00:00:00 2001 From: haosdent Date: Fri, 22 May 2026 14:24:18 +0800 Subject: [PATCH] [CI] Fix test_lora_with_spec_decode on V2 model runner (#43314) Signed-off-by: haosdent Co-authored-by: Wentao Ye <44945378+yewentao256@users.noreply.github.com> --- .../spec_decode/test_lora_with_spec_decode.py | 71 ++++++++++--------- vllm/v1/worker/gpu/model_runner.py | 14 +++- 2 files changed, 51 insertions(+), 34 deletions(-) diff --git a/tests/v1/e2e/spec_decode/test_lora_with_spec_decode.py b/tests/v1/e2e/spec_decode/test_lora_with_spec_decode.py index 9e000223c14..5decfbc0086 100644 --- a/tests/v1/e2e/spec_decode/test_lora_with_spec_decode.py +++ b/tests/v1/e2e/spec_decode/test_lora_with_spec_decode.py @@ -67,6 +67,12 @@ def test_batch_inference_correctness( method, model_name, spec_model_name, lora_path, tp_size = model_setup + prompts = [LORA_TEST_PROMPT_MAP[lora_path]] * 100 + lora_request = LoRARequest("adapter", 1, lora_path) + sampling_params = SamplingParams( + temperature=0.0, top_p=1.0, top_k=-1, seed=SEED, max_tokens=128 + ) + # without speculative decoding ref_llm = LLM( model=model_name, @@ -79,19 +85,14 @@ def test_batch_inference_correctness( max_cpu_loras=1, max_lora_rank=16, ) - - prompts = [LORA_TEST_PROMPT_MAP[lora_path]] * 100 - lora_request = LoRARequest("adapter", 1, lora_path) - sampling_params = SamplingParams( - temperature=0.0, top_p=1.0, top_k=-1, seed=SEED, max_tokens=128 - ) - - ref_outputs = ref_llm.generate( - prompts, sampling_params, lora_request=lora_request - ) - del ref_llm - torch.accelerator.empty_cache() - cleanup_dist_env_and_memory() + try: + ref_outputs = ref_llm.generate( + prompts, sampling_params, lora_request=lora_request + ) + finally: + del ref_llm + torch.accelerator.empty_cache() + cleanup_dist_env_and_memory() lora_spec_llm = LLM( model=model_name, @@ -110,25 +111,29 @@ def test_batch_inference_correctness( max_cpu_loras=1, max_lora_rank=16, ) + try: + lora_spec_outputs = lora_spec_llm.generate( + prompts, sampling_params, lora_request=lora_request + ) - lora_spec_outputs = lora_spec_llm.generate( - prompts, sampling_params, lora_request=lora_request - ) + matches = 0 + for ref_output, spec_output in zip(ref_outputs, lora_spec_outputs): + if ref_output.outputs[0].text == spec_output.outputs[0].text: + matches += 1 + else: + print(f"ref_output: {ref_output.outputs[0].text}") + print(f"spec_output: {spec_output.outputs[0].text}") - matches = 0 - misses = 0 - for ref_output, spec_output in zip(ref_outputs, lora_spec_outputs): - if ref_output.outputs[0].text == spec_output.outputs[0].text: - matches += 1 - else: - misses += 1 - print(f"ref_output: {ref_output.outputs[0].text}") - print(f"spec_output: {spec_output.outputs[0].text}") - - # Heuristic: expect at least 90% of the prompts to match exactly - # Upon failure, inspect the outputs to check for inaccuracy. - print(f"match ratio: {matches}/{len(ref_outputs)}") - assert matches > int(0.90 * len(ref_outputs)) - del lora_spec_llm - torch.accelerator.empty_cache() - cleanup_dist_env_and_memory() + # Heuristic threshold: under greedy verification, the spec-decode + # output should equal the non-spec output (modulo FP noise from the + # target's verify-path matmul running at seqlen + # num_speculative_tokens+1 vs 1). 90% leaves slack for that noise. + threshold = int(0.90 * len(ref_outputs)) + print(f"match ratio: {matches}/{len(ref_outputs)}") + assert matches > threshold, ( + f"match ratio {matches}/{len(ref_outputs)} <= {threshold}" + ) + finally: + del lora_spec_llm + torch.accelerator.empty_cache() + cleanup_dist_env_and_memory() diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 15a2e71f59e..3b965dfc6e1 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -1091,13 +1091,25 @@ class GPUModelRunner(LoRAModelRunnerMixin): # FIXME: Replace this with LoRA warmup: # https://github.com/vllm-project/vllm/pull/35536 assert hasattr(self, "lora_manager") - self.lora_manager._adapter_manager.set_adapter_mapping( + adapter_manager = self.lora_manager._adapter_manager + adapter_manager.set_adapter_mapping( LoRAMapping( index_mapping=(0,) * input_batch.num_tokens_after_padding, prompt_mapping=(0,) * input_batch.num_reqs, is_prefill=True, ) ) + seen_wrappers: set[int] = set() + for punica_wrapper in adapter_manager.punica_wrapper_mapping.values(): + if id(punica_wrapper) in seen_wrappers: + continue + seen_wrappers.add(id(punica_wrapper)) + for kernel_meta in ( + punica_wrapper.token_mapping_meta, # type: ignore[attr-defined] + punica_wrapper.prompt_mapping_meta, # type: ignore[attr-defined] + ): + kernel_meta.no_lora_flag_cpu[0] = False + kernel_meta.num_active_loras_cpu[0] = 1 attn_metadata = None slot_mappings_by_layer = None