Skip to content
Draft
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
95 changes: 92 additions & 3 deletions python/freetoken/engine/cache_budget.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING

from freetoken.utils import div_ceil
Expand All @@ -23,9 +24,97 @@ def expert_bytes_per_slot(sources: dict[str, "list[torch.Tensor]"]) -> int:
"""
# marlin/b12x gate_up/down alpha scales are fixed [L*E] residency (do not scale
# with cache_size), so they are intentionally excluded from the per-slot growth term.
# tensor[0].numel() is the per-row element count (one expert slot); see the matching
# slot-byte idiom in kvcache/linear_state_pool.py and kvcache/dsv4_paged_pool.py.
return sum(t[0][0].numel() * t[0].element_size() for t in sources.values())
# The GPU slot cache has one stride per bank, chosen from that bank's largest layer
# row. See the matching slot-byte calculation in OffloadMoeCache.
return sum(
max(layer[0].numel() * layer.element_size() for layer in per_layer)
for per_layer in sources.values()
)


@dataclass(frozen=True)
class GeometryPoolPlan:
layer_ids: tuple[int, ...]
row_bytes: tuple[int, ...]
slots: int


def plan_geometry_pool_slots(
row_bytes_by_layer: list[tuple[int, ...]],
*,
legacy_cache_size: int,
num_experts: int,
top_k: int,
max_decode_batch: int,
) -> tuple[GeometryPoolPlan, ...] | None:
"""Partition fixed max-stride bank arenas into exact-geometry decode pools.

``legacy_cache_size`` remains the external budget denomination. Each bank owns
``legacy_cache_size * max(layer_row_bytes)`` bytes; every planned class must fit
all bank constraints independently.
"""
if not row_bytes_by_layer:
return ()
num_banks = len(row_bytes_by_layer[0])
if num_banks == 0 or any(len(row) != num_banks for row in row_bytes_by_layer):
raise ValueError("every layer must describe the same non-empty bank set")
if (
legacy_cache_size <= 0
or num_experts <= 0
or top_k <= 0
or max_decode_batch <= 0
):
raise ValueError(
"cache size, experts, top_k, and decode batch must be positive"
)

grouped: dict[tuple[int, ...], list[int]] = {}
for layer_id, rows in enumerate(row_bytes_by_layer):
if any(value <= 0 for value in rows):
raise ValueError("geometry row bytes must be positive")
grouped.setdefault(tuple(rows), []).append(layer_id)
classes = [(rows, tuple(layer_ids)) for rows, layer_ids in grouped.items()]
budgets = [
legacy_cache_size * max(rows[bank] for rows in row_bytes_by_layer)
for bank in range(num_banks)
]
floor = min(num_experts, top_k * max_decode_batch)
slots = [floor] * len(classes)

def used(bank: int) -> int:
return sum(slots[i] * classes[i][0][bank] for i in range(len(classes)))

if any(used(bank) > budgets[bank] for bank in range(num_banks)):
return None

targets = [
min(len(layer_ids) * num_experts, len(layer_ids) * top_k * max_decode_batch)
for _, layer_ids in classes
]
caps = [len(layer_ids) * num_experts for _, layer_ids in classes]

def affordable(index: int) -> bool:
rows = classes[index][0]
return all(
used(bank) + rows[bank] <= budgets[bank] for bank in range(num_banks)
)

def fill(limits: list[int]) -> None:
while True:
candidates = [
i for i in range(len(classes)) if slots[i] < limits[i] and affordable(i)
]
if not candidates:
return
index = min(candidates, key=lambda i: (slots[i] / limits[i], i))
slots[index] += 1

fill(targets)
fill(caps)
return tuple(
GeometryPoolPlan(layer_ids=layer_ids, row_bytes=rows, slots=slots[index])
for index, (rows, layer_ids) in enumerate(classes)
)


def net_cache_budget_bytes(
Expand Down
15 changes: 14 additions & 1 deletion python/freetoken/engine/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -613,6 +613,14 @@ def _init_offload_moe_cache(self, config: EngineConfig) -> OffloadMoeCache:
quant_format=banks.quant_format,
decode_target=decode_target,
hybrid_max_fetch=config.moe_hybrid_max_fetch,
geometry_pool_top_k=getattr(
config.model_config, "num_experts_per_tok", 0
),
geometry_pool_max_batch=max(
config.max_running_req,
config.cuda_graph_max_bs or 0,
1,
),
)
# before set_bank_sources: the residency validation and the copy plan's skip of non-pinned layers key on the CPU-layer set
cache.cpu_layer_ids = cpu_layer_ids
Expand Down Expand Up @@ -1375,7 +1383,12 @@ def override(attr: str, value: Any): # this is dangerous, use with caution
from freetoken.moe.cpu_executor import compiled_extension_supports

_act = getattr(model_config, "hidden_act", "silu")
if not _cpu_moe_act_ok:
if bench_fmt == "gguf":
logger.info_rank0(
"benchbw profile recommends hybrid, but GGUF experts do not have a "
"CPU executor; staying on offload"
)
elif not _cpu_moe_act_ok:
logger.info_rank0(
f"benchbw profile recommends hybrid, but the CPU MoE executor does not "
f"support this model's expert activation "
Expand Down
22 changes: 21 additions & 1 deletion python/freetoken/kernel/csrc/gguf/gguf_kernel.cu
Original file line number Diff line number Diff line change
Expand Up @@ -545,7 +545,8 @@ torch::Tensor ggml_moe_a8_vec(
int64_t top_k,
int64_t type,
int64_t row,
int64_t tokens) {
int64_t tokens,
int64_t expert_stride_bytes) {
int col = X.sizes()[1];
const int padded = (col + 512 - 1) / 512 * 512;
const at::cuda::OptionalCUDAGuard device_guard(device_of(X));
Expand All @@ -568,6 +569,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 3:
Expand All @@ -581,6 +583,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 6:
Expand All @@ -594,6 +597,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 7:
Expand All @@ -607,6 +611,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 8:
Expand All @@ -620,6 +625,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 10:
Expand All @@ -633,6 +639,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 11:
Expand All @@ -646,6 +653,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 12:
Expand All @@ -659,6 +667,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 13:
Expand All @@ -672,6 +681,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 14:
Expand All @@ -685,6 +695,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 16:
Expand All @@ -698,6 +709,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 17:
Expand All @@ -711,6 +723,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 18:
Expand All @@ -724,6 +737,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 19:
Expand All @@ -737,6 +751,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 20:
Expand All @@ -750,6 +765,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 21:
Expand All @@ -763,6 +779,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 22:
Expand All @@ -776,6 +793,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 23:
Expand All @@ -789,6 +807,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
case 29:
Expand All @@ -802,6 +821,7 @@ torch::Tensor ggml_moe_a8_vec(
col,
row,
quant_X.stride(0),
expert_stride_bytes,
stream);
break;
}
Expand Down
Loading