test_spec_decode.py 19.7 KB
Newer Older
1
# SPDX-License-Identifier: Apache-2.0
2
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
3
import random
4
from typing import Any
5

6
import pytest
zhiweiz's avatar
zhiweiz committed
7
import torch
8

9
from tests.utils import get_attn_backend_list_based_on_platform, large_gpu_mark
10
from vllm import LLM, SamplingParams
11
12
from vllm.assets.base import VLLM_S3_BUCKET_URL
from vllm.assets.image import VLM_IMAGES_DIR
zhiweiz's avatar
zhiweiz committed
13
from vllm.distributed import cleanup_dist_env_and_memory
14
from vllm.platforms import current_platform
15

16
17
MTP_SIMILARITY_RATE = 0.8

18

19
def get_test_prompts(mm_enabled: bool):
20
    prompt_types = ["repeat", "sentence"]
21
22
    if mm_enabled:
        prompt_types.append("mm")
23
24
25
26
27
    num_prompts = 100
    prompts = []

    random.seed(0)
    random_prompt_type_choices = random.choices(prompt_types, k=num_prompts)
28
    print(f"Prompt types: {random_prompt_type_choices}")
29
30
31
32
33
34

    # Generate a mixed batch of prompts, some of which can be easily
    # predicted by n-gram matching and some which likely cannot.
    for kind in random_prompt_type_choices:
        word_choices = ["test", "temp", "hello", "where"]
        word = random.choice(word_choices)
35
        prompt: str | list[dict[str, Any]] = ""
36
37
38
39
40
41
42
43
44
45
46
47
        if kind == "repeat":
            prompt = f"""
            please repeat the word '{word}' 10 times.
            give no other output than the word at least ten times in a row,
            in lowercase with spaces between each word and without quotes.
            """
        elif kind == "sentence":
            prompt = f"""
            please give a ten-word sentence that
            uses the word {word} at least once.
            give no other output than that simple sentence without quotes.
            """
48
        elif kind == "mm":
49
50
51
52
53
54
55
56
            placeholders = [
                {
                    "type": "image_url",
                    "image_url": {
                        "url": f"{VLLM_S3_BUCKET_URL}/{VLM_IMAGES_DIR}/stop_sign.jpg"
                    },
                }
            ]
57
58
            prompt = [
                *placeholders,
59
                {"type": "text", "text": "The meaning of the image is"},
60
            ]
61
62
63
64
65
        else:
            raise ValueError(f"Unknown prompt type: {kind}")
        prompts.append([{"role": "user", "content": prompt}])

    return prompts
66
67
68
69


@pytest.fixture
def sampling_config():
70
    return SamplingParams(temperature=0, max_tokens=10, ignore_eos=False)
71
72
73
74


@pytest.fixture
def model_name():
75
    return "meta-llama/Llama-3.1-8B-Instruct"
76
77


78
79
80
81
82
83
84
85
@pytest.fixture(autouse=True)
def reset_torch_dynamo():
    """Reset torch dynamo cache before each test"""
    yield
    # Cleanup after test
    torch._dynamo.reset()


86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
@pytest.mark.parametrize(
    "speculative_config",
    [
        {
            "method": "ngram",
            "prompt_lookup_max": 5,
            "prompt_lookup_min": 3,
            "num_speculative_tokens": 3,
        },
        {
            "method": "suffix",
            "suffix_decoding_max_spec_factor": 2.0,
        },
    ],
)
def test_ngram_and_suffix_correctness(
    speculative_config: dict,
103
104
105
106
    monkeypatch: pytest.MonkeyPatch,
    sampling_config: SamplingParams,
    model_name: str,
):
107
    """
108
    Compare the outputs of an original LLM and a speculative LLM
109
    should be the same when using ngram speculative decoding.
110
    """
111
112
113
114
115
116
117
118
119
120
    test_prompts = get_test_prompts(mm_enabled=False)

    ref_llm = LLM(model=model_name, max_model_len=1024)
    ref_outputs = ref_llm.chat(test_prompts, sampling_config)
    del ref_llm
    torch.cuda.empty_cache()
    cleanup_dist_env_and_memory()

    spec_llm = LLM(
        model=model_name,
121
        speculative_config=speculative_config,
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
        max_model_len=1024,
    )
    spec_outputs = spec_llm.chat(test_prompts, sampling_config)
    matches = 0
    misses = 0
    for ref_output, spec_output in zip(ref_outputs, 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 66% of the prompts to match exactly
    # Upon failure, inspect the outputs to check for inaccuracy.
    assert matches >= int(0.66 * len(ref_outputs))
    del spec_llm
    torch.cuda.empty_cache()
    cleanup_dist_env_and_memory()


143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
def test_suffix_decoding_acceptance(
    monkeypatch: pytest.MonkeyPatch,
    sampling_config: SamplingParams,
    model_name: str,
):
    """
    Check that suffix decoding caching takes effect and improves acceptance
    lengths and acceptance rates over multiple runs of the same prompts.
    """
    test_prompts = get_test_prompts(mm_enabled=False)

    spec_llm = LLM(
        model=model_name,
        speculative_config={
            "method": "suffix",
            "suffix_decoding_max_spec_factor": 2.0,
            "suffix_decoding_max_cached_requests": 1000,
        },
        max_model_len=1024,
        disable_log_stats=False,
    )

    # Run several times and check that the accepted tokens increase.
    num_draft = []
    num_accept = []
    for i in range(10):  # Run multiple times to warm up the cache.
        spec_llm.chat(test_prompts, sampling_config)
        # Collect draft and acceptance stats.
        metrics = spec_llm.get_metrics()
        for metric in metrics:
            if metric.name == "vllm:spec_decode_num_draft_tokens":
                num_draft.append(metric.value)
            if metric.name == "vllm:spec_decode_num_accepted_tokens":
                num_accept.append(metric.value)

    # Calculate the acceptance rates for the first and last runs.
    first_accept_tokens = num_accept[0]
    first_draft_tokens = num_draft[0]
    first_accept_rate = first_accept_tokens / first_draft_tokens

    # Take the diff since the stats are cumulative.
    last_accept_tokens = num_accept[-1] - num_accept[-2]
    last_draft_tokens = num_draft[-1] - num_draft[-2]
    last_accept_rate = last_accept_tokens / last_draft_tokens

    # Expect the acceptance length to improve.
    assert first_accept_tokens < last_accept_tokens

    # Expect the acceptance rate to improve.
    assert first_accept_rate < last_accept_rate

194
195
    # Heuristic: expect at least 80.0% acceptance rate at the end.
    assert last_accept_rate > 0.80
196
197
198
199
200
201

    del spec_llm
    torch.cuda.empty_cache()
    cleanup_dist_env_and_memory()


202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
@pytest.mark.parametrize(
    "model_path",
    [
        "RedHatAI/Llama-3.1-8B-Instruct-speculator.eagle3",
        "RedHatAI/Qwen3-8B-speculator.eagle3",
    ],
    ids=["llama3_eagle3_speculator", "qwen3_eagle3_speculator"],
)
def test_speculators_model_integration(
    monkeypatch: pytest.MonkeyPatch,
    sampling_config: SamplingParams,
    model_path: str,
):
    """
    Test that speculators models work with the simplified integration.

    This verifies the `vllm serve <speculator-model>` use case where
    speculative config is automatically detected from the model config
    without requiring explicit --speculative-config argument.

    Tests:
    1. Speculator model is correctly detected
    2. Verifier model is extracted from speculator config
    3. Speculative decoding is automatically enabled
    4. Text generation works correctly
    5. Output matches reference (non-speculative) generation
    """
    monkeypatch.setenv("VLLM_ALLOW_INSECURE_SERIALIZATION", "1")

    # Generate test prompts
    test_prompts = get_test_prompts(mm_enabled=False)

    # First run: Direct speculator model (simplified integration)
    spec_llm = LLM(model=model_path, max_model_len=1024)
    spec_outputs = spec_llm.chat(test_prompts, sampling_config)

    # Verify speculative config was auto-detected
    assert spec_llm.llm_engine.vllm_config.speculative_config is not None, (
        f"Speculative config should be auto-detected for {model_path}"
    )

    spec_config = spec_llm.llm_engine.vllm_config.speculative_config
    assert spec_config.num_speculative_tokens > 0, (
        f"Expected positive speculative tokens, "
        f"got {spec_config.num_speculative_tokens}"
    )

    # Verify draft model is set to the speculator model
    assert spec_config.model == model_path, (
        f"Draft model should be {model_path}, got {spec_config.model}"
    )

    # Extract verifier model for reference run
    verifier_model = spec_llm.llm_engine.vllm_config.model_config.model

    del spec_llm
    torch.cuda.empty_cache()
    cleanup_dist_env_and_memory()

    # Second run: Reference without speculative decoding
    ref_llm = LLM(model=verifier_model, max_model_len=1024)
    ref_outputs = ref_llm.chat(test_prompts, sampling_config)
    del ref_llm
    torch.cuda.empty_cache()
    cleanup_dist_env_and_memory()

    # Compare outputs
    matches = sum(
        1
        for ref, spec in zip(ref_outputs, spec_outputs)
        if ref.outputs[0].text == spec.outputs[0].text
    )

    # Heuristic: expect at least 66% of prompts to match exactly
    assert matches >= int(0.66 * len(ref_outputs)), (
        f"Only {matches}/{len(ref_outputs)} outputs matched. "
        f"Expected at least {int(0.66 * len(ref_outputs))} matches."
    )


282
@pytest.mark.parametrize(
283
    ["model_setup", "mm_enabled", "enable_chunked_prefill"],
284
    [
285
        (("eagle3", "Qwen/Qwen3-8B", "AngelSlim/Qwen3-8B_eagle3", 1), False, False),
286
287
288
289
290
291
292
293
294
295
296
297
298
        pytest.param(
            (
                "eagle3",
                "Qwen/Qwen3-VL-8B-Instruct",
                "taobao-mnn/Qwen3-VL-8B-Instruct-Eagle3",
                1,
            ),
            False,
            False,
            marks=pytest.mark.skip(
                reason="architecture of its eagle3 is LlamaForCausalLMEagle3"
            ),
        ),
299
300
301
302
303
304
305
306
        pytest.param(
            (
                "eagle3",
                "Qwen/Qwen2.5-VL-7B-Instruct",
                "Rayzl/qwen2.5-vl-7b-eagle3-sgl",
                1,
            ),
            False,
307
            False,
308
309
310
311
            marks=pytest.mark.skip(
                reason="Skipping due to its head_dim not being a a multiple of 32"
            ),
        ),
312
        pytest.param(
313
314
315
316
317
318
319
            (
                "eagle",
                "meta-llama/Llama-3.1-8B-Instruct",
                "yuhuili/EAGLE-LLaMA3.1-Instruct-8B",
                1,
            ),
            False,
320
321
322
            True,
            marks=large_gpu_mark(min_gb=40),
        ),  # works on 4x H100
323
324
325
326
327
328
329
330
        (
            (
                "eagle3",
                "meta-llama/Llama-3.1-8B-Instruct",
                "yuhuili/EAGLE3-LLaMA3.1-Instruct-8B",
                1,
            ),
            False,
331
            False,
332
333
334
335
336
337
338
339
340
        ),
        pytest.param(
            (
                "eagle",
                "meta-llama/Llama-4-Scout-17B-16E-Instruct",
                "morgendave/EAGLE-Llama-4-Scout-17B-16E-Instruct",
                4,
            ),
            False,
341
            False,
342
343
344
345
346
347
348
349
350
351
            marks=large_gpu_mark(min_gb=80),
        ),  # works on 4x H100
        pytest.param(
            (
                "eagle",
                "meta-llama/Llama-4-Scout-17B-16E-Instruct",
                "morgendave/EAGLE-Llama-4-Scout-17B-16E-Instruct",
                4,
            ),
            True,
352
            True,
353
354
355
356
357
358
359
360
361
362
            marks=large_gpu_mark(min_gb=80),
        ),  # works on 4x H100
        (
            (
                "eagle",
                "eagle618/deepseek-v3-random",
                "eagle618/eagle-deepseek-v3-random",
                1,
            ),
            False,
363
            False,
364
        ),
365
366
    ],
    ids=[
367
        "qwen3_eagle3",
368
        "qwen3_vl_eagle3",
369
370
371
372
373
374
375
376
377
        "qwen2_5_vl_eagle3",
        "llama3_eagle",
        "llama3_eagle3",
        "llama4_eagle",
        "llama4_eagle_mm",
        "deepseek_eagle",
    ],
)
@pytest.mark.parametrize("attn_backend", get_attn_backend_list_based_on_platform())
378
379
380
def test_eagle_correctness(
    monkeypatch: pytest.MonkeyPatch,
    sampling_config: SamplingParams,
zhiweiz's avatar
zhiweiz committed
381
    model_setup: tuple[str, str, str, int],
382
    mm_enabled: bool,
383
    enable_chunked_prefill: bool,
384
    attn_backend: str,
385
):
386
387
388
389
    if attn_backend == "TREE_ATTN":
        # TODO: Fix this flaky test
        pytest.skip(
            "TREE_ATTN is flaky in the test disable for now until it can be "
390
391
            "resolved (see https://github.com/vllm-project/vllm/issues/22922)"
        )
392

393
394
    # Generate test prompts inside the function instead of using fixture
    test_prompts = get_test_prompts(mm_enabled)
395
    """
396
397
    Compare the outputs of a original LLM and a speculative LLM
    should be the same when using eagle speculative decoding.
zhiweiz's avatar
zhiweiz committed
398
    model_setup: (method, model_name, eagle_model_name, tp_size)
399
    """
400
    with monkeypatch.context() as m:
401
402
403
404
        if "Llama-4-Scout" in model_setup[1] and attn_backend == "FLASH_ATTN":
            # Scout requires default backend selection
            # because vision encoder has head_dim 88 being incompatible
            #  with FLASH_ATTN and needs to fall back to Flex Attn
405
406
407
408
409

            # pass if not ROCm
            if current_platform.is_rocm():
                # TODO: Enable Flex Attn for spec_decode on ROCm
                pytest.skip("Flex Attn for spec_decode not supported on ROCm currently")
410
411
412
        else:
            m.setenv("VLLM_MLA_DISABLE", "1")
            m.setenv("VLLM_ATTENTION_BACKEND", attn_backend)
413

414
415
416
417
418
        if attn_backend == "TRITON_ATTN" and not current_platform.is_rocm():
            pytest.skip(
                "TRITON_ATTN does not support "
                "multi-token eagle spec decode on current platform"
            )
419

420
        if attn_backend == "ROCM_AITER_FA" and current_platform.is_rocm():
421
            if "deepseek" in model_setup[1].lower():
422
                pytest.skip("ROCM_AITER_FA for deepseek not supported on ROCm platform")
423
424
            else:
                m.setenv("VLLM_ROCM_USE_AITER", "1")
425

zhiweiz's avatar
zhiweiz committed
426
        method, model_name, spec_model_name, tp_size = model_setup
427
        max_model_len = 2048
428
        max_num_batched_tokens = 128 if enable_chunked_prefill else max_model_len
429

430
        ref_llm = LLM(
431
            model=model_name, max_model_len=max_model_len, tensor_parallel_size=tp_size
432
        )
433
        ref_outputs = ref_llm.chat(test_prompts, sampling_config)
434
        del ref_llm
zhiweiz's avatar
zhiweiz committed
435
436
        torch.cuda.empty_cache()
        cleanup_dist_env_and_memory()
437

438
439
        spec_llm = LLM(
            model=model_name,
440
            trust_remote_code=True,
zhiweiz's avatar
zhiweiz committed
441
            tensor_parallel_size=tp_size,
442
            speculative_config={
zhiweiz's avatar
zhiweiz committed
443
                "method": method,
444
                "model": spec_model_name,
445
                "num_speculative_tokens": 3,
446
                "max_model_len": max_model_len,
447
            },
448
449
            max_model_len=max_model_len,
            max_num_batched_tokens=max_num_batched_tokens,
450
            enable_chunked_prefill=enable_chunked_prefill,
451
        )
452
453
454
        spec_outputs = spec_llm.chat(test_prompts, sampling_config)
        matches = 0
        misses = 0
455
        for ref_output, spec_output in zip(ref_outputs, spec_outputs):
456
457
458
459
460
461
462
            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}")

463
        # Heuristic: expect at least 60% of the prompts to match exactly
464
        # Upon failure, inspect the outputs to check for inaccuracy.
465
        assert matches > int(0.6 * len(ref_outputs))
466
        del spec_llm
zhiweiz's avatar
zhiweiz committed
467
468
469
470
        torch.cuda.empty_cache()
        cleanup_dist_env_and_memory()


471
472
473
474
475
476
477
478
@pytest.mark.parametrize(
    ["model_setup", "mm_enabled"],
    [
        (("mtp", "XiaomiMiMo/MiMo-7B-Base", 1), False),
        (("mtp", "ZixiQi/DeepSeek-V3-4layers-MTP-FP8", 1), False),
    ],
    ids=["mimo", "deepseek"],
)
479
def test_mtp_correctness(
480
481
    monkeypatch: pytest.MonkeyPatch,
    sampling_config: SamplingParams,
482
    model_setup: tuple[str, str, int],
483
    mm_enabled: bool,
484
):
485
486
    # Generate test prompts inside the function instead of using fixture
    test_prompts = get_test_prompts(mm_enabled)
487
    """
488
    Compare the outputs of a original LLM and a speculative LLM
489
490
    should be the same when using MTP speculative decoding.
    model_setup: (method, model_name, tp_size)
491
    """
492
    with monkeypatch.context() as m:
493
        m.setenv("VLLM_MLA_DISABLE", "1")
494

495
        method, model_name, tp_size = model_setup
496

497
498
499
500
501
502
        ref_llm = LLM(
            model=model_name,
            max_model_len=2048,
            tensor_parallel_size=tp_size,
            trust_remote_code=True,
        )
503
504
        ref_outputs = ref_llm.chat(test_prompts, sampling_config)
        del ref_llm
zhiweiz's avatar
zhiweiz committed
505
506
        torch.cuda.empty_cache()
        cleanup_dist_env_and_memory()
507
508
509

        spec_llm = LLM(
            model=model_name,
510
            trust_remote_code=True,
zhiweiz's avatar
zhiweiz committed
511
            tensor_parallel_size=tp_size,
512
            speculative_config={
zhiweiz's avatar
zhiweiz committed
513
                "method": method,
514
                "num_speculative_tokens": 1,
515
                "max_model_len": 2048,
516
            },
517
            max_model_len=2048,
518
519
520
521
522
523
524
525
526
527
528
529
        )
        spec_outputs = spec_llm.chat(test_prompts, sampling_config)
        matches = 0
        misses = 0
        for ref_output, spec_output in zip(ref_outputs, 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}")

530
        # Heuristic: expect at least 80% of the prompts to match exactly
531
        # Upon failure, inspect the outputs to check for inaccuracy.
532
        assert matches > int(MTP_SIMILARITY_RATE * len(ref_outputs))
533
        del spec_llm
zhiweiz's avatar
zhiweiz committed
534
535
        torch.cuda.empty_cache()
        cleanup_dist_env_and_memory()
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598


@pytest.mark.parametrize(["model_setup", "mm_enabled"], [
    (("mtp", "XiaomiMiMo/MiMo-7B-Base", 1), False),
    (("mtp", "ZixiQi/DeepSeek-V3-4layers-MTP-FP8", 1), False),
],
                         ids=["mimo", "deepseek"])
def test_mtp_correctness(
    monkeypatch: pytest.MonkeyPatch,
    sampling_config: SamplingParams,
    model_setup: tuple[str, str, int],
    mm_enabled: bool,
):
    # Generate test prompts inside the function instead of using fixture
    test_prompts = get_test_prompts(mm_enabled)
    '''
    Compare the outputs of a original LLM and a speculative LLM
    should be the same when using MTP speculative decoding.
    model_setup: (method, model_name, tp_size)
    '''
    with monkeypatch.context() as m:
        m.setenv("VLLM_USE_V1", "1")
        m.setenv("VLLM_MLA_DISABLE", "1")

        method, model_name, tp_size = model_setup

        ref_llm = LLM(model=model_name,
                      max_model_len=2048,
                      tensor_parallel_size=tp_size,
                      trust_remote_code=True)
        ref_outputs = ref_llm.chat(test_prompts, sampling_config)
        del ref_llm
        torch.cuda.empty_cache()
        cleanup_dist_env_and_memory()

        spec_llm = LLM(
            model=model_name,
            trust_remote_code=True,
            tensor_parallel_size=tp_size,
            speculative_config={
                "method": method,
                "num_speculative_tokens": 1,
                "max_model_len": 2048,
            },
            max_model_len=2048,
        )
        spec_outputs = spec_llm.chat(test_prompts, sampling_config)
        matches = 0
        misses = 0
        for ref_output, spec_output in zip(ref_outputs, 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 80% of the prompts to match exactly
        # Upon failure, inspect the outputs to check for inaccuracy.
        assert matches > int(MTP_SIMILARITY_RATE * len(ref_outputs))
        del spec_llm
        torch.cuda.empty_cache()
        cleanup_dist_env_and_memory()