Studio: simplify the inference orchestrator and worker (#6439)

This commit is contained in:
oobabooga 2026-06-18 17:23:16 -03:00 committed by GitHub
commit b4fa48e1b3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 262 additions and 448 deletions

View file

@ -18,7 +18,6 @@ import atexit
import base64
import os
import signal
import structlog
from loggers import get_logger
import multiprocessing as mp
import queue
@ -61,7 +60,6 @@ class InferenceOrchestrator:
self._cmd_queue: Any = None
self._resp_queue: Any = None
self._cancel_event: Any = None # mp.Event — set to cancel generation
self._lock = threading.Lock()
self._gen_lock = threading.Lock() # Serializes generation
# Dispatcher state for compare mode (adapter-controlled requests):
@ -76,7 +74,6 @@ class InferenceOrchestrator:
self.active_model_name: Optional[str] = None
self.models: dict = {}
self.loading_models: set = set()
self.loaded_local_models: list = []
from core.inference.defaults import get_default_models
self._static_models = get_default_models()
@ -84,8 +81,6 @@ class InferenceOrchestrator:
self._top_hub_cache: Optional[list[str]] = None
self._top_models_ready = threading.Event()
self._current_transformers_major: Optional[str] = None # "4" or "5"
atexit.register(self._cleanup)
logger.info("InferenceOrchestrator initialized (subprocess mode)")
@ -392,6 +387,108 @@ class InferenceOrchestrator:
return
logger.warning("Timed out waiting for gen_done after cancel")
# ------------------------------------------------------------------
# Generation command + token-stream helpers (shared by all paths)
# ------------------------------------------------------------------
def _build_generate_cmd(
self,
request_id: str,
image_b64: Optional[str],
*,
messages: list = None,
system_prompt: str = "",
temperature: float = 0.7,
top_p: float = 0.9,
top_k: int = 40,
min_p: float = 0.0,
max_new_tokens: int = 256,
repetition_penalty: float = 1.0,
use_adapter = None,
tools: Optional[list] = None,
enable_thinking: Optional[bool] = None,
reasoning_effort: Optional[str] = None,
preserve_thinking: Optional[bool] = None,
) -> dict:
"""Build the 'generate' command shared by the locked and dispatched paths."""
cmd = {
"type": "generate",
"request_id": request_id,
"messages": messages or [],
"system_prompt": system_prompt,
"image_base64": image_b64,
"temperature": temperature,
"top_p": top_p,
"top_k": top_k,
"min_p": min_p,
"max_new_tokens": max_new_tokens,
"repetition_penalty": repetition_penalty,
}
# Only forward template kwargs the caller set, for older worker compat.
if use_adapter is not None:
cmd["use_adapter"] = use_adapter
if tools is not None:
cmd["tools"] = tools
if enable_thinking is not None:
cmd["enable_thinking"] = enable_thinking
if reasoning_effort is not None:
cmd["reasoning_effort"] = reasoning_effort
if preserve_thinking is not None:
cmd["preserve_thinking"] = preserve_thinking
return cmd
def _consume_token_stream(
self,
read_one,
drain_on_cancel,
*,
crash_context: str,
cancel_event = None,
stats_holder: Optional[dict] = None,
read_timeout: float = 30.0,
) -> Generator[str, None, None]:
"""Yield tokens from a response stream until gen_done/gen_error.
``read_one(timeout)`` returns the next response (or None on timeout) and
owns the queue choice the shared resp_queue under _gen_lock, or a
per-request mailbox on the dispatcher path so this loop stays agnostic
of which queue is read. On cancel, ``drain_on_cancel()`` consumes the
cancel ack from that same source so stale events don't leak into the
next request.
"""
while True:
resp = read_one(read_timeout)
if resp is None:
# Check subprocess health
if not self._ensure_subprocess_alive():
yield f"Error: {self._subprocess_crash_message(crash_context)}"
return
continue
rtype = resp.get("type", "")
if rtype == "status":
continue
# Subprocess-level error (no request_id); request-scoped failures
# arrive as gen_error below.
if rtype == "error" and not resp.get("request_id"):
yield f"Error: {resp.get('error', 'Unknown error')}"
return
if rtype == "token":
# Cancel from route (e.g. SSE connection closed).
if cancel_event is not None and cancel_event.is_set():
self._cancel_generation()
drain_on_cancel()
return
yield resp.get("text", "")
elif rtype == "gen_done":
if stats_holder is not None:
stats_holder["stats"] = resp.get("stats")
return
elif rtype == "gen_error":
yield f"Error: {resp.get('error', 'Unknown error')}"
return
# ------------------------------------------------------------------
# Dispatcher — per-request mailbox routing for compare mode
# ------------------------------------------------------------------
@ -454,13 +551,12 @@ class InferenceOrchestrator:
continue
# No matching mailbox (a _gen_lock reader or orphaned). Can't
# un-get from mp.Queue, so just log.
if rtype not in ("status",):
logger.debug(
"Dispatcher: no mailbox for request_id=%s type=%s, dropping",
rid,
rtype,
)
# un-get from mp.Queue, so just log. (status was handled above.)
logger.debug(
"Dispatcher: no mailbox for request_id=%s type=%s, dropping",
rid,
rtype,
)
def _generate_dispatched(
self,
@ -505,30 +601,23 @@ class InferenceOrchestrator:
if image is not None:
image_b64 = self._pil_to_base64(image)
cmd = {
"type": "generate",
"request_id": request_id,
"messages": messages or [],
"system_prompt": system_prompt,
"image_base64": image_b64,
"temperature": temperature,
"top_p": top_p,
"top_k": top_k,
"min_p": min_p,
"max_new_tokens": max_new_tokens,
"repetition_penalty": repetition_penalty,
}
if use_adapter is not None:
cmd["use_adapter"] = use_adapter
if tools is not None:
cmd["tools"] = tools
if enable_thinking is not None:
cmd["enable_thinking"] = enable_thinking
if reasoning_effort is not None:
cmd["reasoning_effort"] = reasoning_effort
if preserve_thinking is not None:
cmd["preserve_thinking"] = preserve_thinking
cmd = self._build_generate_cmd(
request_id,
image_b64,
messages = messages,
system_prompt = system_prompt,
temperature = temperature,
top_p = top_p,
top_k = top_k,
min_p = min_p,
max_new_tokens = max_new_tokens,
repetition_penalty = repetition_penalty,
use_adapter = use_adapter,
tools = tools,
enable_thinking = enable_thinking,
reasoning_effort = reasoning_effort,
preserve_thinking = preserve_thinking,
)
# Create mailbox BEFORE sending command
mailbox: queue.Queue = queue.Queue()
@ -543,36 +632,22 @@ class InferenceOrchestrator:
yield f"Error: {exc}"
return
# Read tokens from our private mailbox
def read_mailbox(timeout):
try:
return mailbox.get(timeout = timeout)
except queue.Empty:
return None
# Read tokens from our private mailbox (the dispatcher owns resp_queue).
try:
while True:
try:
resp = mailbox.get(timeout = _DISPATCH_READ_TIMEOUT)
except queue.Empty:
# Timeout — check subprocess health
if not self._ensure_subprocess_alive():
yield f"Error: {self._subprocess_crash_message('generation')}"
return
continue
rtype = resp.get("type", "")
if rtype == "token":
# Cancel from route (e.g. SSE connection closed)
if cancel_event is not None and cancel_event.is_set():
self._cancel_generation()
self._drain_mailbox(mailbox, timeout = 5.0)
return
yield resp.get("text", "")
elif rtype == "gen_done":
if stats_holder is not None:
stats_holder["stats"] = resp.get("stats")
return
elif rtype == "gen_error":
yield f"Error: {resp.get('error', 'Unknown error')}"
return
yield from self._consume_token_stream(
read_mailbox,
lambda: self._drain_mailbox(mailbox, timeout = 5.0),
crash_context = "generation",
cancel_event = cancel_event,
stats_holder = stats_holder,
read_timeout = _DISPATCH_READ_TIMEOUT,
)
finally:
with self._mailbox_lock:
self._mailboxes.pop(request_id, None)
@ -721,7 +796,6 @@ class InferenceOrchestrator:
)
if resp.get("success"):
self._current_transformers_major = needed_major
model_info = resp.get("model_info", {})
self.active_model_name = model_info.get("identifier", model_name)
self.models[self.active_model_name] = {
@ -990,127 +1064,42 @@ class InferenceOrchestrator:
self._wait_dispatcher_idle()
# Serialize generation: two concurrent readers on resp_queue would
# consume and drop each other's token events.
# consume and drop each other's token events. Hold _gen_lock across the
# cmd build + send + whole stream so we stay the sole resp_queue reader.
with self._gen_lock:
yield from self._generate_locked(
request_id = str(uuid.uuid4())
image_b64 = self._pil_to_base64(image) if image is not None else None
cmd = self._build_generate_cmd(
request_id,
image_b64,
messages = messages,
system_prompt = system_prompt,
image = image,
temperature = temperature,
top_p = top_p,
top_k = top_k,
min_p = min_p,
max_new_tokens = max_new_tokens,
repetition_penalty = repetition_penalty,
cancel_event = cancel_event,
use_adapter = use_adapter,
tools = tools,
enable_thinking = enable_thinking,
reasoning_effort = reasoning_effort,
preserve_thinking = preserve_thinking,
stats_holder = stats_holder,
)
def _generate_locked(
self,
messages: list = None,
system_prompt: str = "",
image = None,
temperature: float = 0.7,
top_p: float = 0.9,
top_k: int = 40,
min_p: float = 0.0,
max_new_tokens: int = 256,
repetition_penalty: float = 1.0,
cancel_event = None,
use_adapter = None,
tools: Optional[list] = None,
enable_thinking: Optional[bool] = None,
reasoning_effort: Optional[str] = None,
preserve_thinking: Optional[bool] = None,
stats_holder: Optional[dict] = None,
) -> Generator[str, None, None]:
"""Actual generation logic — must be called under _gen_lock."""
request_id = str(uuid.uuid4())
# Convert PIL Image to base64 if needed
image_b64 = None
if image is not None:
image_b64 = self._pil_to_base64(image)
cmd = {
"type": "generate",
"request_id": request_id,
"messages": messages or [],
"system_prompt": system_prompt,
"image_base64": image_b64,
"temperature": temperature,
"top_p": top_p,
"top_k": top_k,
"min_p": min_p,
"max_new_tokens": max_new_tokens,
"repetition_penalty": repetition_penalty,
}
if use_adapter is not None:
cmd["use_adapter"] = use_adapter
# Only forward template kwargs the caller set, for older worker compat.
if tools is not None:
cmd["tools"] = tools
if enable_thinking is not None:
cmd["enable_thinking"] = enable_thinking
if reasoning_effort is not None:
cmd["reasoning_effort"] = reasoning_effort
if preserve_thinking is not None:
cmd["preserve_thinking"] = preserve_thinking
try:
self._send_cmd(cmd)
except RuntimeError as exc:
yield f"Error: {exc}"
return
# We are the only resp_queue reader (under _gen_lock).
while True:
resp = self._read_resp(timeout = 30.0)
if resp is None:
# Check subprocess health
if not self._ensure_subprocess_alive():
yield f"Error: {self._subprocess_crash_message('generation')}"
return
continue
rtype = resp.get("type", "")
# Status messages — skip
if rtype == "status":
continue
# Error without request_id = subprocess-level error
resp_rid = resp.get("request_id")
if rtype == "error" and not resp_rid:
yield f"Error: {resp.get('error', 'Unknown error')}"
try:
self._send_cmd(cmd)
except RuntimeError as exc:
yield f"Error: {exc}"
return
if rtype == "token":
# Cancel from route (e.g. SSE connection closed)
if cancel_event is not None and cancel_event.is_set():
self._cancel_generation()
# Wait for the cancel ack so stale events don't leak into
# the next request.
self._drain_until_gen_done(timeout = 5.0)
return
yield resp.get("text", "")
elif rtype == "gen_done":
if stats_holder is not None:
stats_holder["stats"] = resp.get("stats")
return
elif rtype == "gen_error":
yield f"Error: {resp.get('error', 'Unknown error')}"
return
yield from self._consume_token_stream(
self._read_resp,
lambda: self._drain_until_gen_done(timeout = 5.0),
crash_context = "generation",
cancel_event = cancel_event,
stats_holder = stats_holder,
)
def reset_generation_state(self):
"""Cancel any ongoing generation and reset state."""
@ -1146,8 +1135,6 @@ class InferenceOrchestrator:
if not self.active_model_name:
raise RuntimeError("No active model")
import uuid
request_id = str(uuid.uuid4())
cmd = {
@ -1260,8 +1247,6 @@ class InferenceOrchestrator:
return
with self._gen_lock:
import uuid
request_id = str(uuid.uuid4())
# numpy array -> list for mp.Queue serialization
@ -1290,38 +1275,12 @@ class InferenceOrchestrator:
yield f"Error: {exc}"
return
# Yield tokens — same pattern as _generate_locked
while True:
resp = self._read_resp(timeout = 30.0)
if resp is None:
if not self._ensure_subprocess_alive():
yield ("Error: " + self._subprocess_crash_message("audio input generation"))
return
continue
rtype = resp.get("type", "")
if rtype == "status":
continue
if rtype == "error" and not resp.get("request_id"):
yield f"Error: {resp.get('error', 'Unknown error')}"
return
if rtype == "token":
if cancel_event is not None and cancel_event.is_set():
self._cancel_generation()
self._drain_until_gen_done(timeout = 5.0)
return
yield resp.get("text", "")
elif rtype == "gen_done":
return
elif rtype == "gen_error":
yield f"Error: {resp.get('error', 'Unknown error')}"
return
yield from self._consume_token_stream(
self._read_resp,
lambda: self._drain_until_gen_done(timeout = 5.0),
crash_context = "audio input generation",
cancel_event = cancel_event,
)
# ------------------------------------------------------------------
# Local helpers (no subprocess needed)

View file

@ -13,12 +13,10 @@ mp.Queue, and exits on shutdown or unload. Pattern follows core/training/worker.
from __future__ import annotations
import base64
import structlog
from loggers import get_logger
import os
import queue as _queue
import sys
import threading
import time
import traceback
from io import BytesIO
@ -28,13 +26,19 @@ from typing import Any
logger = get_logger(__name__)
from utils.hardware import apply_gpu_ids
# studio/backend root, prepended to sys.path so the spawned subprocess can
# import the utils/core packages.
_BACKEND_PATH = str(Path(__file__).resolve().parent.parent.parent)
def _ensure_backend_on_path() -> None:
if _BACKEND_PATH not in sys.path:
sys.path.insert(0, _BACKEND_PATH)
def _activate_transformers_version(model_name: str) -> None:
"""Activate the correct transformers version BEFORE any ML imports."""
# Ensure backend is on path for utils imports.
backend_path = str(Path(__file__).resolve().parent.parent.parent)
if backend_path not in sys.path:
sys.path.insert(0, backend_path)
_ensure_backend_on_path()
from utils.transformers_version import activate_transformers_for_subprocess
@ -63,160 +67,97 @@ def _resize_image(img, max_size: int = 800):
def _send_response(resp_queue: Any, response: dict) -> None:
"""Send a response to the parent process."""
"""Send a response to the parent process; stamps ``ts`` if absent."""
response.setdefault("ts", time.time())
try:
resp_queue.put(response)
except (OSError, ValueError) as exc:
logger.error("Failed to send response: %s", exc)
def _clean_token(value: str | None) -> str | None:
"""Normalize an HF token: blank or whitespace-only becomes None."""
return value if value and value.strip() else None
def _build_model_config(config: dict):
"""Build a ModelConfig from the config dict."""
from utils.models import ModelConfig
model_name = config["model_name"]
hf_token = config.get("hf_token")
hf_token = hf_token if hf_token and hf_token.strip() else None
gguf_variant = config.get("gguf_variant")
mc = ModelConfig.from_identifier(
model_id = model_name,
hf_token = hf_token,
gguf_variant = gguf_variant,
hf_token = _clean_token(config.get("hf_token")),
gguf_variant = config.get("gguf_variant"),
)
if not mc:
raise ValueError(f"Invalid model identifier: {model_name}")
return mc
def _get_hf_download_state(model_names: list[str] | None = None) -> tuple[int, bool] | None:
"""Return (total_bytes, has_incomplete) for the HF Hub cache, or None on error.
_NEMOTRON_TRUST_SUBSTRINGS = ("nemotron_h", "nemotron-h", "nemotron-3-nano")
With *model_names*, only those models' ``blobs/`` dirs are checked (faster);
accepts multiple names so LoRA loads can watch adapter + base repos at once.
*has_incomplete* is True when any ``*.incomplete`` files exist (download
active). None means state could not be determined, so callers skip stall logic.
def _needs_nemotron_trust(model_name: str, hf_token: str | None = None) -> bool:
"""Whether *model_name* is a NemotronH/Nano model that needs trust_remote_code.
NemotronH/Nano have config-parsing bugs that require it. Must NOT match
Llama-Nemotron (standard Llama arch), so also require the unsloth/ or nvidia/
namespace, and a genuine first-party Hub repo (not a local path or a spoof
name starting with "unsloth/"). The repo check is authenticated so private
first-party repos still resolve, and runs only after the cheap checks pass.
"""
mn = model_name.lower()
if not (
any(sub in mn for sub in _NEMOTRON_TRUST_SUBSTRINGS)
and (mn.startswith("unsloth/") or mn.startswith("nvidia/"))
):
return False
from utils.security.trusted_org import is_trusted_org_repo
return is_trusted_org_repo(model_name, hf_token = hf_token)
def _resolve_lora_4bit(mc, load_in_4bit: bool) -> bool:
"""Reconcile load_in_4bit with a LoRA adapter's recorded training method.
lora -> base is full precision (4bit off); qlora -> base is quantized (4bit
on); unknown method -> force off only when the base is not a -bnb-4bit repo.
A missing or unreadable adapter_config.json leaves the value unchanged.
"""
if not (mc.is_lora and mc.path):
return load_in_4bit
adapter_cfg_path = Path(mc.path) / "adapter_config.json"
if not adapter_cfg_path.exists():
return load_in_4bit
import json
try:
from huggingface_hub.constants import HF_HUB_CACHE
cache = Path(HF_HUB_CACHE)
if not cache.exists():
return (0, False)
total = 0
has_incomplete = False
blobs_dirs: list[Path] = []
if model_names:
from utils.paths import resolve_cached_repo_id_case
for name in model_names:
if not name:
continue
# Skip local filesystem paths -- HF IDs (org/model) never start
# with / . ~ or contain backslashes.
if name.startswith(("/", ".", "~")) or "\\" in name:
continue
name = resolve_cached_repo_id_case(name)
# HF cache dir format: models--org--name (slashes -> --).
cache_dir_name = "models--" + name.replace("/", "--")
blobs_dir = cache / cache_dir_name / "blobs"
if blobs_dir.exists():
blobs_dirs.append(blobs_dir)
else:
blobs_dirs = list(cache.glob("models--*/blobs"))
for bdir in blobs_dirs:
for f in bdir.iterdir():
try:
if f.is_file():
total += f.stat().st_size
if f.name.endswith(".incomplete"):
has_incomplete = True
except OSError:
pass
return (total, has_incomplete)
except Exception as e:
logger.debug("Failed to determine HF download state: %s", e)
return None
def _start_heartbeat(
resp_queue: Any,
interval: float = 30.0,
stall_timeout: float = 180.0,
xet_disabled: bool = False,
model_names: list[str] | None = None,
) -> threading.Event:
"""Start a daemon thread that sends periodic status heartbeats.
A stall is reported only when ``*.incomplete`` files are present (download
active) AND cache size hasn't changed for *stall_timeout* seconds. When the
download finishes the timer resets, so post-download init (quantization, GPU
weight load) isn't misclassified as a stall. Returns a stop event.
"""
stop = threading.Event()
transport = "https" if xet_disabled else "xet"
def _beat():
state = _get_hf_download_state(model_names)
last_size = state[0] if state is not None else 0
last_change = time.monotonic()
while not stop.wait(interval):
state = _get_hf_download_state(model_names)
now = time.monotonic()
# Skip stall logic if we cannot measure the cache.
if state is None:
_send_response(
resp_queue,
{
"type": "status",
"message": f"Loading model ({transport} transport)...",
"ts": time.time(),
},
)
continue
current_size, has_incomplete = state
if current_size != last_size:
last_size = current_size
last_change = now
# Only fire stall while .incomplete files confirm an active download;
# reset the timer otherwise so model init isn't counted as a stall.
if not has_incomplete:
last_change = now
elif now - last_change >= stall_timeout:
_send_response(
resp_queue,
{
"type": "stall",
"message": (
f"Download appears stalled ({transport} transport) "
f"-- no progress for {int(now - last_change)}s"
),
"ts": time.time(),
},
)
# fire once -- the orchestrator will kill us
return
_send_response(
resp_queue,
{
"type": "status",
"message": f"Loading model ({transport} transport)...",
"ts": time.time(),
},
with open(adapter_cfg_path) as f:
adapter_cfg = json.load(f)
training_method = adapter_cfg.get("unsloth_training_method")
if training_method == "lora" and load_in_4bit:
logger.info("adapter_config.json says lora — setting load_in_4bit=False")
return False
if training_method == "qlora" and not load_in_4bit:
logger.info("adapter_config.json says qlora — setting load_in_4bit=True")
return True
if (
not training_method
and mc.base_model
and "-bnb-4bit" not in mc.base_model.lower()
and load_in_4bit
):
logger.info(
"No training method, base model has no -bnb-4bit — setting load_in_4bit=False"
)
t = threading.Thread(target = _beat, daemon = True)
t.start()
return stop
return False
except Exception as e:
logger.warning("Could not read adapter_config.json: %s", e)
return load_in_4bit
def _handle_load(backend, config: dict, resp_queue: Any) -> None:
@ -224,61 +165,15 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
try:
mc = _build_model_config(config)
hf_token = config.get("hf_token")
hf_token = hf_token if hf_token and hf_token.strip() else None
hf_token = _clean_token(config.get("hf_token"))
load_in_4bit = _resolve_lora_4bit(mc, config.get("load_in_4bit", True))
# Auto-detect quantization for LoRA adapters.
load_in_4bit = config.get("load_in_4bit", True)
if mc.is_lora and mc.path:
import json
from pathlib import Path
adapter_cfg_path = Path(mc.path) / "adapter_config.json"
if adapter_cfg_path.exists():
try:
with open(adapter_cfg_path) as f:
adapter_cfg = json.load(f)
training_method = adapter_cfg.get("unsloth_training_method")
if training_method == "lora" and load_in_4bit:
logger.info("adapter_config.json says lora — setting load_in_4bit=False")
load_in_4bit = False
elif training_method == "qlora" and not load_in_4bit:
logger.info("adapter_config.json says qlora — setting load_in_4bit=True")
load_in_4bit = True
elif not training_method:
if (
mc.base_model
and "-bnb-4bit" not in mc.base_model.lower()
and load_in_4bit
):
logger.info(
"No training method, base model has no -bnb-4bit — setting load_in_4bit=False"
)
load_in_4bit = False
except Exception as e:
logger.warning("Could not read adapter_config.json: %s", e)
# Auto-enable trust_remote_code only for NemotronH/Nano (config parsing
# bugs require it). Must NOT match Llama-Nemotron (standard Llama arch).
from utils.security.trusted_org import is_trusted_org_repo
_NEMOTRON_TRUST_SUBSTRINGS = ("nemotron_h", "nemotron-h", "nemotron-3-nano")
trust_remote_code = config.get("trust_remote_code", False)
if not trust_remote_code:
model_name = config["model_name"]
_mn_lower = model_name.lower()
if (
any(sub in _mn_lower for sub in _NEMOTRON_TRUST_SUBSTRINGS)
and (_mn_lower.startswith("unsloth/") or _mn_lower.startswith("nvidia/"))
# Genuine first-party Hub repo only (not a local/spoof name starting
# with "unsloth/"); authenticated so private repos resolve.
and is_trusted_org_repo(model_name, hf_token = hf_token)
):
trust_remote_code = True
logger.info(
"Auto-enabled trust_remote_code for Nemotron model: %s",
model_name,
)
if not trust_remote_code and _needs_nemotron_trust(config["model_name"], hf_token = hf_token):
trust_remote_code = True
logger.info(
"Auto-enabled trust_remote_code for Nemotron model: %s", config["model_name"]
)
# Malware gate: a poisoned pickle deserializes during from_pretrained even
# with trust_remote_code False, so check HF's security scan (metadata-only)
@ -301,7 +196,6 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
"message": _fs.reason,
"error_kind": "malware_blocked",
"security": _fs.response_payload(),
"ts": time.time(),
},
)
return
@ -334,25 +228,26 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
),
"error_kind": "remote_code_blocked",
"remote_code": _rc.response_payload(),
"ts": time.time(),
},
)
return
# Heartbeat every 30s so the orchestrator knows we're alive during slow loads.
xet_disabled = os.environ.get("HF_HUB_DISABLE_XET") == "1"
# Heartbeat keeps the orchestrator's inactivity deadline alive during slow
# loads; a no-progress Xet download is reported as a stall so the parent
# can respawn over HTTP. Watch model + base repos (base is the LoRA
# download bottleneck).
from utils.hf_xet_fallback import start_watchdog
# Watch model + base repos (base download is the LoRA bottleneck).
watch_repos = [mc.identifier]
base = getattr(mc, "base_model", None)
if base and str(base) != mc.identifier:
watch_repos.append(str(base))
heartbeat_stop = _start_heartbeat(
resp_queue,
interval = 30.0,
xet_disabled = xet_disabled,
model_names = watch_repos,
heartbeat_stop = start_watchdog(
repo_ids = watch_repos,
on_stall = lambda msg: _send_response(resp_queue, {"type": "stall", "message": msg}),
on_heartbeat = lambda msg: _send_response(resp_queue, {"type": "status", "message": msg}),
xet_disabled = os.environ.get("HF_HUB_DISABLE_XET") == "1",
)
try:
success = backend.load_model(
@ -367,7 +262,6 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
heartbeat_stop.set()
if success:
# Build model_info for the parent to mirror.
model_info = {
"identifier": mc.identifier,
"display_name": mc.display_name,
@ -380,13 +274,11 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
"audio_type": getattr(mc, "audio_type", None),
"has_audio_input": getattr(mc, "has_audio_input", False),
}
_bm = getattr(backend, "models", {}) or {}
_entry = (
_bm.get(mc.identifier) or _bm.get(getattr(backend, "active_model_name", None)) or {}
)
try:
_bm = getattr(backend, "models", {}) or {}
_entry = (
_bm.get(mc.identifier)
or _bm.get(getattr(backend, "active_model_name", None))
or {}
)
_context_length = _entry.get("context_length")
if _context_length is not None:
model_info["context_length"] = int(_context_length)
@ -394,12 +286,6 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
logger.warning("context_length forward failed: %s", _ctx_exc)
# Forward chat_template_info so the parent can classify capabilities.
try:
_bm = getattr(backend, "models", {}) or {}
_entry = (
_bm.get(mc.identifier)
or _bm.get(getattr(backend, "active_model_name", None))
or {}
)
_tpl_info = _entry.get("chat_template_info")
if isinstance(_tpl_info, dict):
model_info["chat_template_info"] = {
@ -417,7 +303,6 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
"type": "loaded",
"success": True,
"model_info": model_info,
"ts": time.time(),
},
)
else:
@ -427,7 +312,6 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
"type": "loaded",
"success": False,
"error": "Failed to load model",
"ts": time.time(),
},
)
@ -439,7 +323,6 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
"success": False,
"error": str(exc),
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)
@ -505,7 +388,6 @@ def _handle_generate(backend, cmd: dict, resp_queue: Any, cancel_event) -> None:
"type": "token",
"request_id": request_id,
"text": cumulative_text,
"ts": time.time(),
},
)
@ -516,7 +398,6 @@ def _handle_generate(backend, cmd: dict, resp_queue: Any, cancel_event) -> None:
"request_id": request_id,
# usage/timings from the MLX backend (None elsewhere).
"stats": getattr(backend, "last_generation_stats", None),
"ts": time.time(),
},
)
logger.info("Finished text generation for request_id=%s", request_id)
@ -530,7 +411,6 @@ def _handle_generate(backend, cmd: dict, resp_queue: Any, cancel_event) -> None:
"request_id": request_id,
"error": str(exc),
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)
@ -559,7 +439,6 @@ def _handle_generate_audio(backend, cmd: dict, resp_queue: Any) -> None:
"request_id": request_id,
"wav_base64": base64.b64encode(wav_bytes).decode("ascii"),
"sample_rate": sample_rate,
"ts": time.time(),
},
)
logger.info("Finished audio generation for request_id=%s", request_id)
@ -573,7 +452,6 @@ def _handle_generate_audio(backend, cmd: dict, resp_queue: Any) -> None:
"request_id": request_id,
"error": str(exc),
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)
@ -622,7 +500,6 @@ def _handle_generate_audio_input(backend, cmd: dict, resp_queue: Any, cancel_eve
"type": "token",
"request_id": request_id,
"text": text_chunk,
"ts": time.time(),
},
)
@ -631,7 +508,6 @@ def _handle_generate_audio_input(backend, cmd: dict, resp_queue: Any, cancel_eve
{
"type": "gen_done",
"request_id": request_id,
"ts": time.time(),
},
)
logger.info("Finished audio input generation for request_id=%s", request_id)
@ -645,7 +521,6 @@ def _handle_generate_audio_input(backend, cmd: dict, resp_queue: Any, cancel_eve
"request_id": request_id,
"error": str(exc),
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)
@ -664,7 +539,6 @@ def _handle_unload(backend, cmd: dict, resp_queue: Any) -> None:
{
"type": "unloaded",
"model_name": model_name,
"ts": time.time(),
},
)
except Exception as exc:
@ -675,7 +549,6 @@ def _handle_unload(backend, cmd: dict, resp_queue: Any) -> None:
"type": "unloaded",
"model_name": model_name,
"error": str(exc),
"ts": time.time(),
},
)
@ -712,9 +585,7 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
model_name = config["model_name"]
# ── 0. MLX fast-path — skip torch/transformers ──
backend_path = str(Path(__file__).resolve().parent.parent.parent)
if backend_path not in sys.path:
sys.path.insert(0, backend_path)
_ensure_backend_on_path()
from utils.hardware import hardware as _hw
@ -737,7 +608,7 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
backend = MLXInferenceBackend()
_send_response(
resp_queue,
{"type": "status", "message": "Loading model...", "ts": time.time()},
{"type": "status", "message": "Loading model..."},
)
_handle_load(backend, config, resp_queue)
except Exception as exc:
@ -747,7 +618,6 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
"type": "error",
"error": f"MLX inference init failed: {exc}",
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)
return
@ -779,7 +649,7 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
elif cmd_type == "reset":
cancel_event.set()
backend.reset_generation_state()
_send_response(resp_queue, {"type": "reset_ack", "ts": time.time()})
_send_response(resp_queue, {"type": "reset_ack"})
elif cmd_type == "status":
_send_response(
resp_queue,
@ -791,7 +661,6 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
for k, v in backend.models.items()
},
"loading": list(backend.loading_models),
"ts": time.time(),
},
)
elif cmd_type == "shutdown":
@ -805,7 +674,6 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
"request_id": cmd.get("request_id"),
"error": str(exc),
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)
return
@ -820,7 +688,6 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
"type": "error",
"error": f"Failed to activate transformers version: {exc}",
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)
return
@ -844,13 +711,10 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
{
"type": "status",
"message": "Importing Unsloth...",
"ts": time.time(),
},
)
backend_path = str(Path(__file__).resolve().parent.parent.parent)
if backend_path not in sys.path:
sys.path.insert(0, backend_path)
_ensure_backend_on_path()
from core.inference.inference import InferenceBackend
@ -865,7 +729,6 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
"type": "error",
"error": f"Failed to import ML libraries: {exc}",
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)
return
@ -879,7 +742,6 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
{
"type": "status",
"message": "Loading model...",
"ts": time.time(),
},
)
@ -892,7 +754,6 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
"type": "error",
"error": f"Failed to initialize inference backend: {exc}",
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)
return
@ -923,7 +784,6 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
_handle_generate(backend, cmd, resp_queue, cancel_event)
elif cmd_type == "load":
# Unload the current model before loading the new one.
if backend.active_model_name:
backend.unload_model(backend.active_model_name)
_handle_load(backend, cmd, resp_queue)
@ -951,7 +811,6 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
resp_queue,
{
"type": "reset_ack",
"ts": time.time(),
},
)
@ -970,22 +829,20 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
for name, info in backend.models.items()
},
"loading": list(backend.loading_models),
"ts": time.time(),
},
)
elif cmd_type == "shutdown":
logger.info("Shutdown command received, exiting")
for model_name in list(backend.models.keys()):
for name in list(backend.models.keys()):
try:
backend.unload_model(model_name)
backend.unload_model(name)
except Exception:
pass
_send_response(
resp_queue,
{
"type": "shutdown_ack",
"ts": time.time(),
},
)
return
@ -997,7 +854,6 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
{
"type": "error",
"error": f"Unknown command type: {cmd_type}",
"ts": time.time(),
},
)
@ -1009,6 +865,5 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
"type": "error",
"error": f"Command '{cmd_type}' failed: {exc}",
"stack": traceback.format_exc(limit = 20),
"ts": time.time(),
},
)