Fix --mem-fraction-static not accounting for EAGLE draft model KV cache (#23862)
This commit is contained in:
@@ -0,0 +1,71 @@
|
||||
"""Unit tests for ModelConfig shape normalization."""
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _make_text_config(**overrides):
|
||||
defaults = dict(
|
||||
architectures=["MixtralForCausalLM"],
|
||||
model_type="mixtral",
|
||||
hidden_size=4096,
|
||||
num_attention_heads=32,
|
||||
num_hidden_layers=2,
|
||||
vocab_size=32000,
|
||||
num_key_value_heads=8,
|
||||
)
|
||||
defaults.update(overrides)
|
||||
return SimpleNamespace(**defaults)
|
||||
|
||||
|
||||
class TestModelConfigShapes(CustomTestCase):
|
||||
def _derive_shapes(self, text_config):
|
||||
model_config = ModelConfig.__new__(ModelConfig)
|
||||
model_config.hf_config = text_config
|
||||
model_config.hf_text_config = text_config
|
||||
model_config._derive_model_shapes()
|
||||
return model_config
|
||||
|
||||
def test_optional_head_dims_default_when_none(self):
|
||||
text_config = _make_text_config(
|
||||
head_dim=None,
|
||||
v_head_dim=None,
|
||||
swa_head_dim=None,
|
||||
swa_v_head_dim=None,
|
||||
)
|
||||
|
||||
model_config = self._derive_shapes(text_config)
|
||||
|
||||
self.assertEqual(model_config.head_dim, 128)
|
||||
self.assertEqual(model_config.v_head_dim, 128)
|
||||
self.assertEqual(model_config.swa_head_dim, 128)
|
||||
self.assertEqual(model_config.swa_v_head_dim, 128)
|
||||
self.assertEqual(text_config.head_dim, 128)
|
||||
self.assertEqual(text_config.v_head_dim, 128)
|
||||
self.assertEqual(text_config.swa_head_dim, 128)
|
||||
self.assertEqual(text_config.swa_v_head_dim, 128)
|
||||
|
||||
def test_explicit_head_dims_are_preserved(self):
|
||||
text_config = _make_text_config(
|
||||
head_dim=128,
|
||||
v_head_dim=96,
|
||||
swa_head_dim=64,
|
||||
swa_v_head_dim=48,
|
||||
)
|
||||
|
||||
model_config = self._derive_shapes(text_config)
|
||||
|
||||
self.assertEqual(model_config.head_dim, 128)
|
||||
self.assertEqual(model_config.v_head_dim, 96)
|
||||
self.assertEqual(model_config.swa_head_dim, 64)
|
||||
self.assertEqual(model_config.swa_v_head_dim, 48)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -83,6 +83,8 @@ def _make_model_runner(
|
||||
mr.server_args = sa
|
||||
|
||||
spec = MagicMock()
|
||||
spec.is_eagle.return_value = False
|
||||
spec.is_standalone.return_value = False
|
||||
spec.is_dflash.return_value = False
|
||||
spec.is_none.return_value = True
|
||||
mr.spec_algorithm = spec
|
||||
@@ -303,6 +305,35 @@ class TestAllSWAConfigurator(unittest.TestCase):
|
||||
self.assertEqual(config.swa_max_total_num_tokens, 500)
|
||||
|
||||
|
||||
class TestEagleConfigurator(unittest.TestCase):
|
||||
"""EAGLE: draft KV cache must be accounted for so total allocation fits in budget."""
|
||||
|
||||
def test_eagle_does_not_exceed_budget(self):
|
||||
"""Total memory (target + draft KV cache) must not exceed available."""
|
||||
available = 10_000_000
|
||||
num_layers = 32
|
||||
eagle_draft_num_layers = 4
|
||||
|
||||
mr = _make_model_runner(num_layers=num_layers)
|
||||
mr.spec_algorithm.is_eagle.return_value = True
|
||||
mr.spec_algorithm.is_standalone.return_value = False
|
||||
mr.spec_algorithm.is_none.return_value = False
|
||||
mr.eagle_draft_num_layers = eagle_draft_num_layers
|
||||
|
||||
with mock_cpu_env():
|
||||
from sglang.srt.model_executor.pool_configurator import (
|
||||
create_memory_pool_configurator,
|
||||
)
|
||||
|
||||
cfg = create_memory_pool_configurator(mr)
|
||||
config = cfg.calculate_pool_sizes(available, 1)
|
||||
|
||||
full_pt = _full_per_token(mr)
|
||||
total_layers = num_layers + eagle_draft_num_layers
|
||||
used = config.max_total_num_tokens * full_pt * total_layers
|
||||
self.assertLessEqual(used, available)
|
||||
|
||||
|
||||
class TestFactory(unittest.TestCase):
|
||||
def test_default_for_non_swa(self):
|
||||
mr = _make_model_runner(is_hybrid_swa=False)
|
||||
|
||||
@@ -8,11 +8,12 @@ slow path (`organize_draft_results`) for num_steps in {1, 2, 3, 4}.
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.speculative.eagle_utils import organize_draft_results
|
||||
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker
|
||||
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
|
||||
from sglang.srt.utils import get_device
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
@@ -57,6 +58,20 @@ def _make_worker(num_steps: int, num_draft_tokens: int):
|
||||
return worker
|
||||
|
||||
|
||||
def _make_backend_factory(decode_backend, draft_extend_backend):
|
||||
class FakeDraftBackendFactory:
|
||||
def __init__(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
def create_decode_backend(self):
|
||||
return decode_backend
|
||||
|
||||
def create_draft_extend_backend(self):
|
||||
return draft_extend_backend
|
||||
|
||||
return FakeDraftBackendFactory
|
||||
|
||||
|
||||
class TestEagleWorkerV2Topk1FastPath(CustomTestCase):
|
||||
def test_fast_path_matches_slow_path(self):
|
||||
bs = 3
|
||||
@@ -93,5 +108,68 @@ class TestEagleWorkerV2Topk1FastPath(CustomTestCase):
|
||||
worker._rebuild_topk1_chain_buffers()
|
||||
|
||||
|
||||
class TestEagleWorkerV2BackendFallback(CustomTestCase):
|
||||
def test_preserves_initialized_backend_when_draft_extend_backend_is_unset(self):
|
||||
worker = object.__new__(EagleDraftWorker)
|
||||
existing_backend = object()
|
||||
decode_backend = object()
|
||||
worker.server_args = SimpleNamespace()
|
||||
worker.draft_runner = SimpleNamespace(attn_backend=existing_backend)
|
||||
worker.topk = 1
|
||||
worker.speculative_num_steps = 2
|
||||
|
||||
with patch(
|
||||
"sglang.srt.speculative.eagle_worker_v2.DraftBackendFactory",
|
||||
_make_backend_factory(decode_backend, None),
|
||||
):
|
||||
worker.init_attention_backend()
|
||||
|
||||
self.assertIs(worker.draft_attn_backend, decode_backend)
|
||||
self.assertIsNone(worker.draft_extend_attn_backend)
|
||||
self.assertIs(worker.draft_runner.draft_attn_backend, decode_backend)
|
||||
self.assertIs(worker.draft_runner.attn_backend, existing_backend)
|
||||
|
||||
def test_uses_draft_extend_backend_when_available(self):
|
||||
worker = object.__new__(EagleDraftWorker)
|
||||
existing_backend = object()
|
||||
decode_backend = object()
|
||||
draft_extend_backend = object()
|
||||
worker.server_args = SimpleNamespace()
|
||||
worker.draft_runner = SimpleNamespace(attn_backend=existing_backend)
|
||||
worker.topk = 1
|
||||
worker.speculative_num_steps = 2
|
||||
|
||||
with patch(
|
||||
"sglang.srt.speculative.eagle_worker_v2.DraftBackendFactory",
|
||||
_make_backend_factory(decode_backend, draft_extend_backend),
|
||||
):
|
||||
worker.init_attention_backend()
|
||||
|
||||
self.assertIs(worker.draft_attn_backend, decode_backend)
|
||||
self.assertIs(worker.draft_extend_attn_backend, draft_extend_backend)
|
||||
self.assertIs(worker.draft_runner.draft_attn_backend, decode_backend)
|
||||
self.assertIs(worker.draft_runner.attn_backend, draft_extend_backend)
|
||||
|
||||
def test_spec_v2_attn_backends_include_draft_extend_fallback(self):
|
||||
target_backend = object()
|
||||
decode_backend = object()
|
||||
fallback_backend = object()
|
||||
|
||||
worker = object.__new__(EAGLEWorkerV2)
|
||||
worker._target_worker = SimpleNamespace(
|
||||
model_runner=SimpleNamespace(attn_backend=target_backend)
|
||||
)
|
||||
worker._draft_worker = SimpleNamespace(
|
||||
draft_attn_backend=decode_backend,
|
||||
draft_extend_attn_backend=None,
|
||||
draft_runner=SimpleNamespace(attn_backend=fallback_backend),
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
worker.spec_v2_attn_backends,
|
||||
(target_backend, decode_backend, fallback_backend),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user