forked from Karylab-cklius/vllm
[BugFix] eagle draft max position embeddings (#49343)
Signed-off-by: JaredforReal <w13431838023@gmail.com> Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com> Co-authored-by: Cyrus Leung <tlleungac@connect.ust.hk>
This commit is contained in:
co-authored by
mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
Cyrus Leung
parent
ad5d29db70
commit
32e657e689
@@ -0,0 +1,115 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Tests for the EAGLE draft ``max_position_embeddings`` override (#48894).
|
||||
|
||||
EAGLE drafts share the target's positional space, but some draft
|
||||
checkpoints (e.g. ``yuhuili/EAGLE3-LLaMA3.1-Instruct-8B``) ship a
|
||||
``max_position_embeddings`` (2048) far smaller than the target's context.
|
||||
That value sizes the draft's rotary ``cos_sin_cache`` while the proposer
|
||||
feeds positions up to the target's ``max_model_len``, so the cache gather
|
||||
goes out of bounds — a device-side assert under torch.compile and silent
|
||||
garbage reads in eager mode. ``SpeculativeConfig`` must raise the draft's
|
||||
value to the target's ``max_model_len``, with a log, for the eagle/eagle3
|
||||
methods only.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
import pytest
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from vllm.config.model import ModelConfig
|
||||
from vllm.config.parallel import ParallelConfig
|
||||
from vllm.config.speculative import SpeculativeConfig
|
||||
|
||||
# All repos are public; only config/tokenizer-config files are fetched.
|
||||
EAGLE3_DRAFT = "yuhuili/EAGLE3-LLaMA3.1-Instruct-8B" # max_position_embeddings=2048
|
||||
LLAMA3_TARGET = "unsloth/Meta-Llama-3.1-8B-Instruct" # max_position_embeddings=131072
|
||||
AR_MODEL = "JackFram/llama-68m" # max_position_embeddings=2048
|
||||
|
||||
_LOGGER = "vllm.config.speculative"
|
||||
_OVERRIDE_MSG = "Overriding draft model max_position_embeddings"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def vllm_caplog(caplog: pytest.LogCaptureFixture, monkeypatch: pytest.MonkeyPatch):
|
||||
"""Make caplog see vLLM logger records (vLLM sets propagate=False)."""
|
||||
monkeypatch.setattr(logging.getLogger("vllm"), "propagate", True)
|
||||
with caplog.at_level(logging.INFO, logger=_LOGGER):
|
||||
yield caplog
|
||||
|
||||
|
||||
def _override_logged(caplog: pytest.LogCaptureFixture) -> bool:
|
||||
return any(_OVERRIDE_MSG in record.getMessage() for record in caplog.records)
|
||||
|
||||
|
||||
@pytest.mark.cpu_test
|
||||
def test_override_raises_smaller_value(vllm_caplog: pytest.LogCaptureFixture):
|
||||
hf_config = PretrainedConfig(max_position_embeddings=2048)
|
||||
SpeculativeConfig._maybe_override_draft_max_position_embeddings(
|
||||
hf_config, target_max_model_len=8192
|
||||
)
|
||||
assert hf_config.max_position_embeddings == 8192
|
||||
assert _override_logged(vllm_caplog)
|
||||
|
||||
|
||||
@pytest.mark.cpu_test
|
||||
def test_override_keeps_sufficient_value(vllm_caplog: pytest.LogCaptureFixture):
|
||||
hf_config = PretrainedConfig(max_position_embeddings=8192)
|
||||
SpeculativeConfig._maybe_override_draft_max_position_embeddings(
|
||||
hf_config, target_max_model_len=8192
|
||||
)
|
||||
assert hf_config.max_position_embeddings == 8192
|
||||
assert not _override_logged(vllm_caplog)
|
||||
|
||||
|
||||
@pytest.mark.cpu_test
|
||||
def test_override_ignores_missing_attribute(vllm_caplog: pytest.LogCaptureFixture):
|
||||
hf_config = PretrainedConfig()
|
||||
hf_config.__dict__.pop("max_position_embeddings", None)
|
||||
SpeculativeConfig._maybe_override_draft_max_position_embeddings(
|
||||
hf_config, target_max_model_len=8192
|
||||
)
|
||||
assert not hasattr(hf_config, "max_position_embeddings")
|
||||
assert not _override_logged(vllm_caplog)
|
||||
|
||||
|
||||
@pytest.mark.cpu_test
|
||||
@pytest.mark.parametrize("method", ["eagle", "eagle3"])
|
||||
def test_eagle_draft_inherits_target_max_model_len(
|
||||
method: str, vllm_caplog: pytest.LogCaptureFixture
|
||||
):
|
||||
target_model_config = ModelConfig(LLAMA3_TARGET)
|
||||
assert target_model_config.max_model_len > 2048
|
||||
speculative_config = SpeculativeConfig(
|
||||
target_model_config=target_model_config,
|
||||
target_parallel_config=ParallelConfig(),
|
||||
model=EAGLE3_DRAFT,
|
||||
method=method,
|
||||
num_speculative_tokens=3,
|
||||
)
|
||||
draft_hf_config = speculative_config.draft_model_config.hf_config
|
||||
assert draft_hf_config.max_position_embeddings == target_model_config.max_model_len
|
||||
assert _override_logged(vllm_caplog)
|
||||
|
||||
|
||||
@pytest.mark.cpu_test
|
||||
def test_independent_draft_model_keeps_its_own_limit(
|
||||
vllm_caplog: pytest.LogCaptureFixture,
|
||||
):
|
||||
"""An independent AR draft may genuinely have a smaller context than the
|
||||
target; its max_position_embeddings must not be resized."""
|
||||
target_model_config = ModelConfig(
|
||||
AR_MODEL, hf_overrides={"max_position_embeddings": 8192}
|
||||
)
|
||||
assert target_model_config.max_model_len == 8192
|
||||
speculative_config = SpeculativeConfig(
|
||||
target_model_config=target_model_config,
|
||||
target_parallel_config=ParallelConfig(),
|
||||
model=AR_MODEL,
|
||||
method="draft_model",
|
||||
num_speculative_tokens=3,
|
||||
)
|
||||
draft_hf_config = speculative_config.draft_model_config.hf_config
|
||||
assert draft_hf_config.max_position_embeddings == 2048
|
||||
assert not _override_logged(vllm_caplog)
|
||||
@@ -0,0 +1,112 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Tests for SpecDecodeBaseProposer.initialize_attn_backend.
|
||||
|
||||
Block tables are stored at kernel-block granularity, so the proposer's
|
||||
``block_size`` (used for slot-mapping math) must be the kernel block size,
|
||||
not the KV cache manager's block size — the two differ when manager blocks
|
||||
are split for the attention kernel. The value must also be deterministic:
|
||||
``_draft_attn_layer_names`` is a set, whose iteration order varies across
|
||||
processes, so anything derived from iteration order must not leak into
|
||||
``block_size``.
|
||||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
import vllm.v1.spec_decode.llm_base_proposer as llm_base_proposer
|
||||
from vllm.v1.spec_decode.eagle import EagleProposer
|
||||
|
||||
SCHEDULER_BLOCK_SIZE = 256
|
||||
KERNEL_BLOCK_SIZE = 64
|
||||
|
||||
|
||||
class _FakeAttentionGroup:
|
||||
def __init__(self, backend, layer_names, kv_cache_spec, kv_cache_group_id):
|
||||
self.backend = backend
|
||||
self.layer_names = list(layer_names)
|
||||
self.kv_cache_spec = kv_cache_spec
|
||||
self.kv_cache_group_id = kv_cache_group_id
|
||||
self.kernel_block_size = None
|
||||
|
||||
def create_metadata_builders(self, vllm_config, device, kernel_block_size=None):
|
||||
self.kernel_block_size = kernel_block_size
|
||||
|
||||
def get_metadata_builder(self):
|
||||
return SimpleNamespace(kv_cache_spec=self.kv_cache_spec)
|
||||
|
||||
|
||||
def _make_proposer(
|
||||
monkeypatch: pytest.MonkeyPatch, layer_names: set[str]
|
||||
) -> EagleProposer:
|
||||
fake_layers = {}
|
||||
for name in layer_names:
|
||||
backend = SimpleNamespace(full_cls_name=lambda: "FakeBackend")
|
||||
fake_layers[name] = SimpleNamespace(
|
||||
get_attn_backend=lambda backend=backend: backend
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
llm_base_proposer, "get_layers_from_vllm_config", lambda *a, **k: fake_layers
|
||||
)
|
||||
monkeypatch.setattr(llm_base_proposer, "AttentionGroup", _FakeAttentionGroup)
|
||||
|
||||
proposer = EagleProposer.__new__(EagleProposer)
|
||||
proposer.vllm_config = None
|
||||
proposer.device = None
|
||||
proposer._draft_attn_layer_names = set(layer_names)
|
||||
proposer.kv_cache_gid = -1
|
||||
proposer.draft_attn_groups = []
|
||||
proposer.block_size = -1
|
||||
return proposer
|
||||
|
||||
|
||||
def _make_kv_cache_config(layer_names: set[str]) -> SimpleNamespace:
|
||||
spec = SimpleNamespace(block_size=SCHEDULER_BLOCK_SIZE)
|
||||
group = SimpleNamespace(layer_names=list(layer_names), kv_cache_spec=spec)
|
||||
return SimpleNamespace(kv_cache_groups=[group])
|
||||
|
||||
|
||||
def test_block_size_uses_kernel_block_size(monkeypatch: pytest.MonkeyPatch):
|
||||
"""The proposer's slot-mapping math runs against the kernel-granularity
|
||||
block table, so block_size must come from kernel_block_sizes."""
|
||||
layer_names = {"draft.0.self_attn.attn"}
|
||||
proposer = _make_proposer(monkeypatch, layer_names)
|
||||
|
||||
proposer.initialize_attn_backend(
|
||||
_make_kv_cache_config(layer_names),
|
||||
kernel_block_sizes=[KERNEL_BLOCK_SIZE],
|
||||
)
|
||||
|
||||
assert proposer.block_size == KERNEL_BLOCK_SIZE
|
||||
assert proposer.block_size != SCHEDULER_BLOCK_SIZE
|
||||
# The metadata builder keeps receiving the kernel block size as well.
|
||||
assert proposer.draft_attn_groups[0].kernel_block_size == KERNEL_BLOCK_SIZE
|
||||
|
||||
|
||||
def test_block_size_falls_back_to_kv_cache_spec(monkeypatch: pytest.MonkeyPatch):
|
||||
layer_names = {"draft.0.self_attn.attn"}
|
||||
proposer = _make_proposer(monkeypatch, layer_names)
|
||||
|
||||
proposer.initialize_attn_backend(
|
||||
_make_kv_cache_config(layer_names), kernel_block_sizes=None
|
||||
)
|
||||
|
||||
assert proposer.block_size == SCHEDULER_BLOCK_SIZE
|
||||
|
||||
|
||||
def test_draft_layer_iteration_is_deterministic(monkeypatch: pytest.MonkeyPatch):
|
||||
"""_draft_attn_layer_names is a set; the attention groups built from it
|
||||
must not depend on its (process-random) iteration order."""
|
||||
layer_names = {"draft.c.attn", "draft.a.attn", "draft.b.attn"}
|
||||
expected_order = sorted(layer_names)
|
||||
|
||||
for insertion_order in (expected_order, expected_order[::-1]):
|
||||
proposer = _make_proposer(monkeypatch, set(insertion_order))
|
||||
proposer.initialize_attn_backend(
|
||||
_make_kv_cache_config(set(insertion_order)),
|
||||
kernel_block_sizes=[KERNEL_BLOCK_SIZE],
|
||||
)
|
||||
assert len(proposer.draft_attn_groups) == 1
|
||||
assert proposer.draft_attn_groups[0].layer_names == expected_order
|
||||
assert proposer.block_size == KERNEL_BLOCK_SIZE
|
||||
@@ -910,6 +910,15 @@ class SpeculativeConfig:
|
||||
f"Unsupported speculative method: '{self.method}'"
|
||||
)
|
||||
|
||||
if self.method in ("eagle", "eagle3"):
|
||||
# EAGLE drafts share the target's positional space; a
|
||||
# draft checkpoint with a smaller max_position_embeddings
|
||||
# than the target under-sizes its rotary cache (#48894).
|
||||
SpeculativeConfig._maybe_override_draft_max_position_embeddings(
|
||||
self.draft_model_config.hf_config,
|
||||
self.target_model_config.max_model_len,
|
||||
)
|
||||
|
||||
# Replace hf_config for EAGLE draft_model
|
||||
if self.method in ("eagle", "eagle3", "dflash"):
|
||||
from vllm.transformers_utils.configs.eagle import EAGLEConfig
|
||||
@@ -1130,6 +1139,40 @@ class SpeculativeConfig:
|
||||
)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _maybe_override_draft_max_position_embeddings(
|
||||
draft_hf_config: PretrainedConfig,
|
||||
target_max_model_len: int,
|
||||
) -> None:
|
||||
"""Raise an EAGLE draft's max_position_embeddings up to the target's.
|
||||
|
||||
The proposer feeds the draft positions up to the target's
|
||||
max_model_len, while max_position_embeddings sizes the draft's
|
||||
rotary cos_sin_cache. A smaller checkpoint value (e.g. 2048 for
|
||||
yuhuili/EAGLE3-LLaMA3.1-Instruct-8B) makes that cache gather go
|
||||
out of bounds (#48894).
|
||||
|
||||
Args:
|
||||
draft_hf_config: The draft model's HF config, mutated in place.
|
||||
target_max_model_len: The target model's max_model_len.
|
||||
"""
|
||||
draft_max_position_embeddings = getattr(
|
||||
draft_hf_config, "max_position_embeddings", None
|
||||
)
|
||||
if (
|
||||
draft_max_position_embeddings is None
|
||||
or draft_max_position_embeddings >= target_max_model_len
|
||||
):
|
||||
return
|
||||
logger.info(
|
||||
"Overriding draft model max_position_embeddings from %d to the "
|
||||
"target model's max_model_len (%d); EAGLE drafts share the "
|
||||
"target's positional space.",
|
||||
draft_max_position_embeddings,
|
||||
target_max_model_len,
|
||||
)
|
||||
draft_hf_config.max_position_embeddings = target_max_model_len
|
||||
|
||||
@staticmethod
|
||||
def _verify_and_get_draft_tp(
|
||||
target_parallel_config: ParallelConfig,
|
||||
|
||||
@@ -1740,7 +1740,10 @@ class SpecDecodeBaseProposer:
|
||||
|
||||
attention_groups: dict[tuple[str, str], AttentionGroup] = {}
|
||||
if kv_cache_spec is not None:
|
||||
for layer_name in self._draft_attn_layer_names:
|
||||
# _draft_attn_layer_names is a set; iterate in sorted order so
|
||||
# that attention_groups (and anything derived from its first
|
||||
# element) is deterministic across processes.
|
||||
for layer_name in sorted(self._draft_attn_layer_names):
|
||||
attn_backend = all_attn_layers[layer_name].get_attn_backend()
|
||||
backend_key = attn_backend.full_cls_name()
|
||||
if backend_key not in attention_groups:
|
||||
@@ -1772,9 +1775,20 @@ class SpecDecodeBaseProposer:
|
||||
attention_groups[backend_key].layer_names.append(layer_name)
|
||||
|
||||
self.draft_attn_groups = list(attention_groups.values())
|
||||
self.block_size = (
|
||||
self.draft_attn_groups[0].get_metadata_builder().kv_cache_spec.block_size
|
||||
)
|
||||
if kernel_block_sizes is not None and 0 <= self.kv_cache_gid < len(
|
||||
kernel_block_sizes
|
||||
):
|
||||
# Slot mappings are computed against the block table, which is
|
||||
# stored at kernel-block granularity. Use the kernel block size
|
||||
# rather than the KV cache manager's block size; the two differ
|
||||
# when manager blocks are split for the attention kernel.
|
||||
self.block_size = kernel_block_sizes[self.kv_cache_gid]
|
||||
else:
|
||||
self.block_size = (
|
||||
self.draft_attn_groups[0]
|
||||
.get_metadata_builder()
|
||||
.kv_cache_spec.block_size
|
||||
)
|
||||
logger.debug("Using block size %d for drafting layers", self.block_size)
|
||||
|
||||
def _determine_batch_execution_and_padding(
|
||||
|
||||
Reference in New Issue
Block a user