Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion python/freetoken/engine/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -490,7 +490,11 @@ def _resolve_auto_moe_cache_size(self, config: EngineConfig, banks) -> tuple[int
num_experts=num_experts,
total_experts=total_experts,
prefill_overlap=config.moe_prefill_overlap,
kv_reserve_tokens=max(config.kv_reserve_tokens, min_reserve),
kv_reserve_tokens=max(
config.kv_reserve_tokens,
min_reserve,
(config.num_page_override or 0) * page_tokens,
),
page_size=page_tokens,
quant_format=banks.quant_format,
)
Expand Down
29 changes: 16 additions & 13 deletions tests/engine/test_cache_budget.py
Original file line number Diff line number Diff line change
Expand Up @@ -265,7 +265,7 @@ class tp_info:
assert fixed == 0


def test_engine_resolve_auto_moe_cache_size_maps_kwargs():
def test_engine_resolve_auto_moe_cache_size_maps_kwargs(monkeypatch):
import torch

from freetoken.engine.engine import Engine
Expand All @@ -292,6 +292,7 @@ class StubConfig:
memory_ratio = 0.9
moe_prefill_overlap = True
kv_reserve_tokens = 0
num_page_override = 64
swa_full_tokens_ratio = 0.2
swa_num_pages_override = None
model_config = StubModelConfig()
Expand All @@ -314,21 +315,23 @@ class StubBanks:
engine._weights_bytes = 1_000_000
engine._pool_cls = MHAKVCache # __init__ skipped -> install the generic pool family

size, pages, overlap = engine._resolve_auto_moe_cache_size(StubConfig(), StubBanks())
captured = {}

# cross-check against the same pure functions, proving the kwarg mapping is faithful
from freetoken.engine.cache_budget import expert_bytes_per_slot, resolve_moe_cache_auto
from freetoken.kvcache.mha_pool import MHAKVCache
def fake_resolve(**kwargs):
captured.update(kwargs)
return 8, 64, True

cache_per_page, fixed, _, _ = MHAKVCache.kv_cost(StubConfig())
expected = resolve_moe_cache_auto(
baseline_free=10_000_000, weights_bytes=1_000_000, memory_ratio=0.9,
cache_per_page=cache_per_page, fixed_cache_size=fixed,
per_expert_bytes=expert_bytes_per_slot(StubBanks.sources),
num_experts=4, total_experts=8, prefill_overlap=True,
kv_reserve_tokens=0, page_size=16, quant_format="bf16",
monkeypatch.setattr(
"freetoken.engine.cache_budget.resolve_moe_cache_auto", fake_resolve
)
assert (size, pages, overlap) == expected
got = engine._resolve_auto_moe_cache_size(StubConfig(), StubBanks())

assert got == (8, 64, True)
assert captured["kv_reserve_tokens"] == 64 * 16
assert captured["page_size"] == 16
assert captured["num_experts"] == 4
assert captured["total_experts"] == 8
assert captured["per_expert_bytes"] == 512 + 256


# ---------------------------------------------------------------------------
Expand Down