Spaces:
Running on Zero
Running on Zero
File size: 15,823 Bytes
9e637cd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 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 194 195 196 197 198 199 200 201 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 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 | """Local inference engine: MLX on Apple Silicon, with a prefix-reusing prompt cache.
Design notes that matter for anyone changing this file:
* **One persistent KV cache per engine, reused across every generation.**
The agent's prompt is dominated by a fixed prefix -- the system prompt plus
the JSON schemas of ~28 tools, which together are several thousand tokens.
Re-prefilling that on every tool step is what made the previous
implementation feel slow (measured: 6.4 s of prefill per step, five to six
steps per question). Here the cache is kept between calls and only the
tokens that actually differ from what the cache already holds are fed to
the model, so the fixed prefix is prefilled exactly once per process.
* **Generated tokens stay in the cache too.** A tool-calling turn is
append-only: prompt, then the assistant's tool call, then the tool result,
then more assistant text. Tracking generated tokens alongside prompt tokens
means continuing that turn costs only the tool-result tokens.
* **Streaming is real.** `stream()` yields text as the model produces it.
Nothing here buffers a whole response and re-emits it word by word.
"""
from __future__ import annotations
import os
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Generator, Iterable, Sequence
PROJECT_ROOT = Path(__file__).resolve().parent.parent
DEFAULT_MODEL = os.environ.get("CONTROLAI_MODEL", "mlx-community/Qwen3-14B-4bit")
# The project's own LoRA adapters are deliberately NOT loaded by default. Both
# of them regressed the behaviour they were meant to improve: `behavior_v1`
# emits a spurious empty `<tool_call></tool_call>` as its first output on
# essentially every prompt (so no tool ever runs), and `sft_v2` generates empty
# output when no tools are exposed and string-typed numbers when they are.
# Set CONTROLAI_ADAPTER=<path> to load one anyway for A/B work.
DEFAULT_ADAPTER = os.environ.get("CONTROLAI_ADAPTER") or None
@dataclass
class SamplingConfig:
"""Decoding parameters. Defaults follow Qwen3's own non-thinking recipe."""
temperature: float = 0.7
top_p: float = 0.8
top_k: int = 20
# Qwen3 recommends presence_penalty over the blunt repetition_penalty that
# the previous implementation applied at 1.15 across the board. A flat
# repetition penalty is actively harmful for this workload: it penalises
# the repeated structural tokens that matrices and JSON are made of
# (`[`, `0`, `,`) exactly when the model is emitting a tool call.
presence_penalty: float = 0.5
max_tokens: int = 1024
def with_(self, **kw: Any) -> "SamplingConfig":
merged = {**self.__dict__, **{k: v for k, v in kw.items() if v is not None}}
return SamplingConfig(**merged)
@dataclass
class Chunk:
"""One streamed piece of model output."""
text: str
token: int
thinking: bool = False
tool_call: bool = False
@dataclass
class Stats:
prompt_tokens: int = 0
cached_tokens: int = 0
generated_tokens: int = 0
prefill_seconds: float = 0.0
decode_seconds: float = 0.0
@property
def decode_tps(self) -> float:
return self.generated_tokens / self.decode_seconds if self.decode_seconds else 0.0
class LocalEngine:
"""Streaming text generation against a locally held MLX model."""
def __init__(
self,
model_id: str = DEFAULT_MODEL,
adapter_path: str | None = DEFAULT_ADAPTER,
sampling: SamplingConfig | None = None,
max_cache_tokens: int = 32768,
) -> None:
from mlx_lm import load
self.model_id = model_id
self.adapter_path = adapter_path
self.sampling = sampling or SamplingConfig()
self.max_cache_tokens = max_cache_tokens
t0 = time.time()
if adapter_path:
self.model, self.tokenizer = load(model_id, adapter_path=adapter_path)
else:
self.model, self.tokenizer = load(model_id)
self.load_seconds = time.time() - t0
self._cache: list[Any] | None = None
self._cache_tokens: list[int] = []
self.last_stats = Stats()
# Resolved once: whether this checkpoint's chat template understands
# Qwen3-style `enable_thinking`, and the ids of the think delimiters.
self.supports_thinking = self._probe_thinking_support()
# `<think>`, `</think>`, `<tool_call>` and `</tool_call>` are each a
# single special token in the Qwen3 vocabulary. Watching for the token
# id rather than matching the rendered string is exact: it cannot be
# defeated by a tag split across two streamed chunks, and it costs one
# integer comparison per token instead of a substring scan.
self._think_open = self._single_token("<think>")
self._think_close = self._single_token("</think>")
self._tool_open = self._single_token("<tool_call>")
self._tool_close = self._single_token("</tool_call>")
# ------------------------------------------------------------------ setup
def _token_ids(self, text: str) -> list[int]:
try:
return self.tokenizer.encode(text, add_special_tokens=False)
except TypeError:
return self.tokenizer.encode(text)
def _single_token(self, text: str) -> int | None:
"""The id of `text` if the tokenizer represents it as one token."""
ids = self._token_ids(text)
return ids[0] if len(ids) == 1 else None
def _probe_thinking_support(self) -> bool:
probe = [{"role": "user", "content": "hi"}]
try:
on = self.tokenizer.apply_chat_template(
probe, tokenize=False, add_generation_prompt=True, enable_thinking=True
)
off = self.tokenizer.apply_chat_template(
probe, tokenize=False, add_generation_prompt=True, enable_thinking=False
)
except Exception:
return False
return on != off
# -------------------------------------------------------------- rendering
def render(
self,
messages: Sequence[dict[str, Any]],
tools: Sequence[dict[str, Any]] | None = None,
enable_thinking: bool = False,
) -> str:
kwargs: dict[str, Any] = {
"tokenize": False,
"add_generation_prompt": True,
}
if tools:
kwargs["tools"] = list(tools)
if self.supports_thinking:
kwargs["enable_thinking"] = enable_thinking
return self.tokenizer.apply_chat_template(list(messages), **kwargs)
def encode(self, text: str) -> list[int]:
return self.tokenizer.encode(text)
def count_tokens(self, text: str) -> int:
return len(self.tokenizer.encode(text))
# ------------------------------------------------------------------ cache
def _align_cache(self, tokens: list[int]) -> list[int]:
"""Point the persistent cache at the longest prefix of `tokens` it
already holds, and return the tokens that still need prefilling."""
from mlx_lm.models.cache import (
can_trim_prompt_cache,
make_prompt_cache,
trim_prompt_cache,
)
reusable = 0
if self._cache is not None:
limit = min(len(self._cache_tokens), len(tokens))
while reusable < limit and self._cache_tokens[reusable] == tokens[reusable]:
reusable += 1
# A cache that cannot be trimmed back to the divergence point is worse
# than no cache: it would silently condition generation on stale
# tokens. Rebuild instead.
if self._cache is not None and reusable < len(self._cache_tokens):
if can_trim_prompt_cache(self._cache):
trim_prompt_cache(self._cache, len(self._cache_tokens) - reusable)
else:
self._cache, reusable = None, 0
if self._cache is None or reusable == 0:
self._cache = make_prompt_cache(self.model)
self._cache_tokens = []
reusable = 0
# MLX must be fed at least one token; an exact cache hit therefore
# rewinds by one and replays the final token.
if reusable == len(tokens) and reusable > 0:
from mlx_lm.models.cache import trim_prompt_cache as _trim
_trim(self._cache, 1)
reusable -= 1
self._cache_tokens = list(tokens[:reusable])
self.last_stats.cached_tokens = reusable
return list(tokens[reusable:])
def reset_cache(self) -> None:
self._cache = None
self._cache_tokens = []
def prewarm(self, text: str) -> int:
"""Prefill a prompt prefix so the first real question doesn't pay for it.
Returns the number of tokens now resident in the cache. Called at
startup with the system-prompt-plus-tool-schemas prefix, which turns
first-question latency from a multi-second prefill into a cache hit.
"""
import mlx.core as mx
from mlx_lm import stream_generate
tokens = self.encode(text)
to_feed = self._align_cache(tokens)
if to_feed:
self.model(mx.array(to_feed)[None], cache=self._cache)
mx.eval([c.state for c in self._cache])
self._cache_tokens = list(tokens)
# Prefilling alone leaves the single-token decode kernels uncompiled,
# so the first real question still paid several seconds of Metal
# warm-up. Generate and discard one token against a throwaway cache to
# force that compilation now, without disturbing the prefix cache.
from mlx_lm.models.cache import make_prompt_cache
scratch = make_prompt_cache(self.model)
for _ in stream_generate(
self.model, self.tokenizer, [tokens[-1]], max_tokens=1, prompt_cache=scratch
):
break
return len(tokens)
# ------------------------------------------------------------- generation
def stream(
self,
prompt: str | list[int],
sampling: SamplingConfig | None = None,
stop: Iterable[str] = (),
think_budget: int | None = None,
) -> Generator[Chunk, None, None]:
"""Yield output chunks as they are generated.
`stop` sequences end generation as soon as they appear (the sequence
itself is emitted, since tool-call parsing wants the closing tag).
`think_budget` caps how many tokens may be spent inside a `<think>`
block: on overrun the block is closed by hand and the model is made to
answer, which bounds worst-case latency on a reasoning model.
"""
from mlx_lm import stream_generate
from mlx_lm.sample_utils import make_logits_processors, make_sampler
cfg = sampling or self.sampling
tokens = self.encode(prompt) if isinstance(prompt, str) else list(prompt)
t0 = time.time()
to_feed = self._align_cache(tokens)
self.last_stats = Stats(
prompt_tokens=len(tokens),
cached_tokens=len(tokens) - len(to_feed),
)
sampler = make_sampler(temp=cfg.temperature, top_p=cfg.top_p, top_k=cfg.top_k)
logits_processors = (
make_logits_processors(presence_penalty=cfg.presence_penalty)
if cfg.presence_penalty
else None
)
stop = tuple(s for s in stop if s)
stop_ids = {self._tool_close} if "</tool_call>" in stop and self._tool_close else set()
text_stops = tuple(s for s in stop if not (s == "</tool_call>" and self._tool_close))
# Only the tail can contain a partial stop sequence, so matching over a
# bounded window keeps this O(1) per token instead of rescanning the
# whole response.
window = max((len(s) for s in text_stops), default=0) + 8
emitted: list[int] = []
tail = ""
budget_left = think_budget
in_think = False
in_tool_call = False
first_token_at: float | None = None
remaining = cfg.max_tokens
while remaining > 0:
forced_close = False
for resp in stream_generate(
self.model,
self.tokenizer,
to_feed,
max_tokens=remaining,
sampler=sampler,
logits_processors=logits_processors,
prompt_cache=self._cache,
):
if first_token_at is None:
first_token_at = time.time()
self.last_stats.prefill_seconds = first_token_at - t0
emitted.append(resp.token)
self._cache_tokens.append(resp.token)
remaining -= 1
text = resp.text
# Exact, token-id state transitions. The tags themselves are
# never emitted -- the caller gets the content and the flags.
if resp.token == self._think_open:
in_think = True
continue
if resp.token == self._think_close:
in_think = False
continue
if resp.token == self._tool_open:
# Unlike the think tags, this one is emitted: the caller
# parses the `<tool_call>...</tool_call>` block out of the
# raw text. The flag lets it suppress the same text from
# the user-visible stream.
in_tool_call = True
if text and not in_think:
tail = (tail + text)[-window:] if window else ""
# Fallback for a checkpoint whose think tags are not single
# tokens; harmless when the ids above already matched.
if self._think_open is None and "<think>" in tail:
in_think = True
if self._think_close is None and "</think>" in tail:
in_think = False
yield Chunk(text=text, token=resp.token, thinking=in_think, tool_call=in_tool_call)
if in_think and budget_left is not None:
budget_left -= 1
if budget_left <= 0:
forced_close = True
break
if resp.token in stop_ids or (text_stops and any(s in tail for s in text_stops)):
remaining = 0
break
else:
remaining = 0
if not forced_close:
break
# Overran the thinking budget: close the block ourselves and let
# the same cache continue straight into the answer.
closer = "\n</think>\n\n"
closer_ids = self._token_ids(closer)
self._cache_tokens.extend(closer_ids)
to_feed = closer_ids
# Deliberately not yielded: this is a control action on the model,
# not model output. Emitting it put a bare "</think>" at the top of
# the answer whenever the budget was reached.
in_think = False
budget_left = None
now = time.time()
self.last_stats.generated_tokens = len(emitted)
self.last_stats.decode_seconds = now - (first_token_at or now)
if len(self._cache_tokens) > self.max_cache_tokens:
self.reset_cache()
def generate(
self,
prompt: str | list[int],
sampling: SamplingConfig | None = None,
stop: Iterable[str] = (),
think_budget: int | None = None,
) -> str:
return "".join(
c.text for c in self.stream(prompt, sampling, stop, think_budget)
)
|