Spaces:
Running on Zero
Running on Zero
File size: 14,912 Bytes
1476b6c e9a64ca 3e5d998 e9a64ca 3e5d998 e9a64ca 1476b6c e384d86 e9a64ca 1476b6c e9a64ca 1476b6c e9a64ca 1476b6c e9a64ca 1476b6c e9a64ca 1476b6c e384d86 1476b6c | 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 | """CUDA inference engine: the same contract as `LocalEngine`, on transformers.
`LocalEngine` is MLX, so it runs on Apple Silicon and nowhere else. The public
demo Space runs on Linux/NVIDIA (ZeroGPU), which needs a second path. This is
that path and nothing more -- it is not a supported way to run ControlAI locally,
and `agent.py`, `registry.py` and everything under `tools/` are untouched by it.
The seam is `ControlAgent(engine=...)`.
It mirrors `LocalEngine`'s design decisions rather than reaching for
`model.generate`:
* **Prefix-reusing KV cache.** `DynamicCache.crop()` trims the cache to the
longest prefix the incoming prompt shares with it, so the ~8k-token
system-prompt-plus-tool-schema prefix is prefilled once per process, not
once per tool step. `model.generate` cannot express that.
* **A hand-written decode loop.** `think_budget` closes an overrunning
`<think>` block by *injecting* the closing token, which no `generate`
callback can do, and `presence_penalty` (not `repetition_penalty`) is
applied for the reason spelled out in `engine.py`: a flat repetition penalty
punishes the `[`, `0`, `,` that matrices and JSON are made of.
**Loading is bf16 and moves to the GPU with an explicit `.to("cuda")`. Do not
reintroduce `device_map` or bitsandbytes, and do not construct this class lazily
at request time.** All three break on ZeroGPU, which is the only place this file
runs, and all three fail the same way:
RuntimeError: Low-level CUDA init (`torch._C._cuda_init`) reached. This
means ZeroGPU's PyTorch CUDA emulation mode did not intercept a CUDA
operation in your code.
ZeroGPU patches torch during the import of the Space's entry module and attaches
real hardware only inside a `@spaces.GPU` call. Only CUDA operations inside that
import window are intercepted, so **where this object is constructed matters as
much as how**: `app_space.py` builds it at module scope for exactly that reason.
`device_map` fails on top of that, because it routes transformers through
`caching_allocator_warmup`, which calls `torch.empty(..., device="cuda")`
directly. bitsandbytes in turn *requires* `device_map`, so 4-bit quantisation is
unavailable here — which is why the model has to be small enough in bf16.
That is why the model is Qwen3-8B rather than the 14B run locally: bf16 8B is
~16GB to download against ~28GB, and a Space rebuild re-downloads from scratch.
"""
from __future__ import annotations
import os
import time
from typing import Any, Generator, Iterable, Sequence
from controlai_agent.engine import Chunk, SamplingConfig, Stats
def dtype_kwarg(dtype: Any) -> dict[str, Any]:
"""`{"dtype": ...}` or `{"torch_dtype": ...}`, whichever this release takes.
transformers renamed the argument in 4.56 and the old spelling is gone in
recent releases. Pinning below 4.56 to keep using it is what broke the Space
build: the platform force-installs gradio 6.x, which requires
huggingface-hub >= 1.16, while every transformers < 4.56 requires < 1.0.
Detecting the spelling costs two lines and pins nothing.
"""
import transformers
major, minor = (int(x) for x in transformers.__version__.split(".")[:2])
key = "dtype" if (major, minor) >= (4, 56) else "torch_dtype"
return {key: dtype}
# Smaller than the 14B run locally, deliberately: see the module docstring.
DEFAULT_TORCH_MODEL = os.environ.get("CONTROLAI_MODEL_TORCH", "Qwen/Qwen3-8B")
class TorchEngine:
"""Streaming generation against a CUDA-resident transformers model."""
def __init__(
self,
model_id: str = DEFAULT_TORCH_MODEL,
adapter_path: str | None = None,
sampling: SamplingConfig | None = None,
max_cache_tokens: int = 32768,
) -> None:
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
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()
self.tokenizer = AutoTokenizer.from_pretrained(model_id)
# No device_map and no quantization_config -- see the module docstring.
# Load to CPU, then move with .to(), which ZeroGPU's emulation intercepts.
# The CPU branch exists so the decode loop can be exercised on a small
# model off a GPU box; it is far too slow to actually serve.
cuda = torch.cuda.is_available()
self.model = AutoModelForCausalLM.from_pretrained(
model_id, **dtype_kwarg(torch.bfloat16 if cuda else torch.float32)
)
if adapter_path:
from peft import PeftModel
self.model = PeftModel.from_pretrained(self.model, adapter_path)
self.model = self.model.to("cuda" if cuda else "cpu")
self.model.eval()
self.load_seconds = time.time() - t0
# Verifiable in the Space logs: if this says cpu on the Space, the .to()
# did not take and every request will be minutes rather than seconds.
print(f"[torch] {model_id} on {next(self.model.parameters()).device} "
f"in {self.load_seconds:.1f}s")
self._cache: Any | None = None
self._cache_tokens: list[int] = []
self.last_stats = Stats()
self.supports_thinking = self._probe_thinking_support()
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:
ids = self._token_ids(text)
return ids[0] if len(ids) == 1 else None
def _probe_thinking_support(self) -> bool:
try:
self.tokenizer.apply_chat_template(
[{"role": "user", "content": "x"}],
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)
return True
except (TypeError, ValueError):
return False
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 reset_cache(self) -> None:
self._cache = None
self._cache_tokens = []
def _align_cache(self, tokens: list[int]) -> list[int]:
"""Trim the cache to the longest prefix it shares with `tokens`.
Returns the suffix that still has to be fed to the model. Mirrors
`LocalEngine._align_cache`; see that docstring for why this exists.
"""
from transformers import DynamicCache
if self._cache is None or not self._cache_tokens:
self._cache = DynamicCache()
self._cache_tokens = []
return list(tokens)
shared = 0
for a, b in zip(self._cache_tokens, tokens):
if a != b:
break
shared += 1
# Never keep the whole prompt: the model needs at least one token to
# run forward on, or there are no logits to sample from.
if shared >= len(tokens):
shared = len(tokens) - 1
if shared > self.max_cache_tokens:
shared = 0
if shared == 0:
self._cache = DynamicCache()
self._cache_tokens = []
return list(tokens)
if shared < len(self._cache_tokens):
# crop() is what makes prefix reuse possible. It has moved around
# between transformers releases and this file no longer pins a
# version, so losing it costs speed, not correctness: fall back to
# re-prefilling the whole prompt.
if not hasattr(self._cache, "crop"):
self._cache = DynamicCache()
self._cache_tokens = []
return list(tokens)
self._cache.crop(shared)
self._cache_tokens = list(tokens[:shared])
return list(tokens[shared:])
def prewarm(self, text: str) -> int:
"""Prefill a prompt prefix so the first real question doesn't pay for it."""
import torch
tokens = self.encode(text)
to_feed = self._align_cache(tokens)
if to_feed:
with torch.inference_mode():
self.model(
input_ids=torch.tensor([to_feed], device=self.model.device),
past_key_values=self._cache,
use_cache=True,
)
self._cache_tokens = list(tokens)
return len(tokens)
# ------------------------------------------------------------- generation
def _sample(self, logits: Any, cfg: SamplingConfig, seen: set[int]) -> int:
import torch
logits = logits.float()
if cfg.presence_penalty and seen:
idx = torch.tensor(sorted(seen), device=logits.device)
logits[idx] -= cfg.presence_penalty
if cfg.temperature <= 0:
return int(torch.argmax(logits).item())
logits = logits / cfg.temperature
if cfg.top_k and cfg.top_k > 0:
kth = torch.topk(logits, min(cfg.top_k, logits.numel())).values[-1]
logits = logits.masked_fill(logits < kth, float("-inf"))
probs = torch.softmax(logits, dim=-1)
if cfg.top_p and 0 < cfg.top_p < 1:
ordered, order = torch.sort(probs, descending=True)
cumulative = torch.cumsum(ordered, dim=-1)
# Keep the first token that crosses top_p, so the mask is never
# empty even when one token already carries more than top_p mass.
drop = cumulative - ordered > cfg.top_p
ordered[drop] = 0.0
probs = torch.zeros_like(probs).scatter_(0, order, ordered)
probs = probs / probs.sum()
return int(torch.multinomial(probs, 1).item())
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. See `LocalEngine.stream`."""
import torch
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),
)
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))
window = max((len(s) for s in text_stops), default=0) + 8
eos_ids = {self.tokenizer.eos_token_id}
for extra in ("<|im_end|>", "<|endoftext|>"):
tid = self._single_token(extra)
if tid is not None:
eos_ids.add(tid)
eos_ids.discard(None)
emitted: list[int] = []
seen: set[int] = set()
tail = ""
thinking = False
think_tokens = 0
prefill_done = False
with torch.inference_mode():
step_input = to_feed
while len(emitted) < cfg.max_tokens:
out = self.model(
input_ids=torch.tensor([step_input], device=self.model.device),
past_key_values=self._cache,
use_cache=True,
)
self._cache = out.past_key_values
self._cache_tokens.extend(step_input)
if not prefill_done:
self.last_stats.prefill_seconds = time.time() - t0
t1 = time.time()
prefill_done = True
token = self._sample(out.logits[0, -1, :], cfg, seen)
if token in eos_ids:
break
seen.add(token)
emitted.append(token)
if token == self._think_open:
thinking, think_tokens = True, 0
elif token == self._think_close:
thinking = False
elif thinking:
think_tokens += 1
text = self.tokenizer.decode([token], skip_special_tokens=False)
yield Chunk(
text=text,
token=token,
thinking=thinking,
tool_call=token == self._tool_open,
)
step_input = [token]
# Overran the reasoning budget: close the block by hand and make
# the model answer. Bounds worst-case latency on a reasoning model.
if (
thinking
and think_budget
and think_tokens >= think_budget
and self._think_close is not None
):
thinking = False
emitted.append(self._think_close)
yield Chunk(text="</think>", token=self._think_close, thinking=False)
step_input = [token, self._think_close]
if token in stop_ids:
break
if text_stops:
tail = (tail + text)[-window:]
if any(s in tail for s in text_stops):
break
self.last_stats.generated_tokens = len(emitted)
self.last_stats.decode_seconds = time.time() - (t1 if prefill_done else t0)
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))
|