unsloth/studio/backend/core/inference
Daniel Han 3bfc83781d
Runtime MTP fallback for tensor parallelism (try MTP, recover if it crashes) (#6324)
* Disable MTP speculative decoding under tensor parallelism

Follow-up to #6040 (Studio tensor-parallel support).

MTP-draft speculative decoding plus --split-mode tensor crashes the CUDA
flash-attn kernel at decode time. The startup /health probe only checks that
llama-server comes up, so the existing MTP-drop fallback (keyed on startup
health) never fires and the server dies on the first generation instead.

Gate MTP off when a tensor attempt actually engages: this runs before the
VRAM planner (so no drafter memory is reserved) and before the speculative
flag build (so no --model-draft / --spec-type is emitted). Ngram modes use no
draft model and are kept, and mtp+ngram degrades to ngram rather than off. The
layer-split fallback re-runs with tensor_parallel False and restores MTP.

The reason is surfaced as spec_fallback_reason "tensor_parallel" so the
settings sheet explains why MTP is off instead of prompting a llama.cpp update.

Verified on unsloth/gemma-4-26B-A4B-it-GGUF:UD-Q4_K_XL across 4x B200: the
load now emits --split-mode tensor with no MTP flags and generation completes
without the prior decode crash.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Make tensor-parallel MTP gate test format-independent

The assertion pinned the multi-line `speculative_type = (` form, but ruff
collapses it onto one line, so match `speculative_type =` instead.

* Recover from MTP+tensor-parallel crashes at runtime instead of banning MTP

MTP-draft speculative decoding under --split-mode tensor usually works, but can
crash llama-server's CUDA flash-attn kernel at decode time (the prompt-cache
checkpoint-restore path). The earlier fix statically disabled MTP whenever
tensor parallelism was on, which is not future-proof and gives up the MTP
speedup even though it normally works.

Replace the static ban with a try/recover, mirroring the existing load-time
MTP-drop fallback:

- Load-time decode probe: after the server passes /health under tensor +
  MTP, run one tiny /completion to exercise the draft path. A failure flips
  the load unhealthy so the existing fallback respawns with --spec-default.
  Catches a hard incompatibility that crashes on the first decode.

- Generation-time recovery: snapshot the load kwargs after a healthy load,
  and if llama-server exits mid-generation while MTP + tensor parallelism were
  active, quietly reload the same model with speculative decoding off (one
  single-flight background reload) and surface spec_fallback_reason=runtime_error.
  Catches the rare mid-generation crash the probe and load-time fallback miss.

No persistent ban: a later fresh load re-tries MTP, so this self-heals if a
future llama.cpp supports the combo. Verified on gemma-4-26B-A4B + 4x B200:
MTP runs normally, and killing llama-server mid-generation reloads it without
MTP and serves the next request cleanly.

* Address review feedback on the MTP runtime fallback

- Authenticate the decode probe: direct-stream mode runs llama-server with
  --api-key, so the unauthenticated /completion probe got a 401 and falsely
  dropped MTP. Attach the same bearer auth the other internal requests use.
- Re-check the cancel flag inside the recovery thread after the death poll,
  so an /unload that races the reload can't resurrect the dropped model.
- Schedule the no-MTP recovery on the connection-error paths it was missing:
  generate_chat_completion's ConnectError branch, the OpenAI passthrough
  typed (RemoteProtocolError/ReadError/CloseError) stream catch, and the
  Anthropic passthrough generic stream catch. Previously a server that died
  before reconnect, or a typed mid-stream error, skipped the reload.

* Cover every request path with the MTP+tensor crash recovery via a watchdog

The runtime MTP-crash recovery only fired from request handlers that
observed the failure, so the direct llama-server proxy endpoints
(/v1/completions, /v1/responses, the OpenAI/Anthropic passthrough
transports) -- and a crash with no request in flight -- could leave a
dead server. Add a single background watchdog, armed only on a healthy
MTP + tensor-parallel load, that polls the subprocess and routes an
unexpected death into the existing single-flight no-MTP reload. It is
stopped inside _kill_process (the one deliberate-termination chokepoint)
so a planned reload/unload is never mistaken for a crash, and re-checks
the stop flag after a detected exit to close the kill-vs-poll race. The
reload turns MTP off, so the replacement server arms no watchdog and the
fallback cannot loop; a later fresh load still re-tries MTP.

* Harden MTP+tensor crash recovery: stale-load race, pass-through MTP, requested mode

Address review findings on the runtime MTP-crash recovery:

- Stale-load race: the recovery thread snapshotted the crashed load, waited up
  to 5s for the process to confirm dead, then only checked the cancel flag
  before replaying load_model. A concurrent user load clears that flag, so the
  stale snapshot could reload the old model over the user's new one. Make the
  load lock re-entrant and run the staleness check (cancel + same process +
  unchanged snapshot) under it, atomically with the reload.

- Pass-through MTP: MTP can also be requested via a user --spec-type in
  extra_args or LLAMA_ARG_SPEC_TYPE, where Studio emits no spec flags and
  _speculative_type stays unset, so the probe/watchdog/recovery never engaged.
  Track _mtp_runtime_fallback_active from the actual launched config and gate on
  it; on the no-MTP reload, append a last-wins --spec-default so the replay drops
  MTP regardless of source (and the load-time fallback does the same).

- Requested mode: the off-reload reset _requested_spec_mode to off, so after a
  status refresh the UI showed a bare Off with the runtime-error note suppressed
  and would not retry MTP. Restore the original requested mode after the reload,
  matching the startup MTP fallback.

- Snapshot the extra_args list by value so a caller mutating it cannot corrupt
  the recovery snapshot.

Tests: test_tensor_parallel.py + test_llama_server_args.py green (303 passed).

* Trim verbose comments in the MTP+tensor crash recovery

Tighten the docstrings and inline comments added for the runtime MTP recovery
(watchdog, probe, reload, gating) to succinct one/two-line forms; no code
change (verified comment-only).

---------

Co-authored-by: danielhanchen <michaelhan2050@gmail.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2026-06-18 05:42:21 -07:00
..
__init__.py Studio: make code comments and docstrings more succinct (#6029) 2026-06-08 23:07:28 -07:00
_html_to_md.py Reduce and tighten code comments and docstrings repo-wide (#6095) 2026-06-08 23:09:51 -07:00
anthropic_compat.py Studio: improve OpenAI- and Anthropic-compatible API spec compliance (#6010) 2026-06-09 17:13:25 +02:00
api_monitor.py Studio: trim serving-log noise and surface llama-server engine stats (#6377) 2026-06-17 05:37:57 -07:00
audio_codecs.py Reduce and tighten code comments and docstrings repo-wide (#6095) 2026-06-08 23:09:51 -07:00
chat_template_helpers.py Studio: make code comments and docstrings more succinct (#6029) 2026-06-08 23:07:28 -07:00
chat_templates.py Studio: bundle Gemma 4 chat templates (E2B/E4B + larger) and auto-apply to unsloth/gemma-4-*-GGUF (#6245) 2026-06-12 05:49:39 -07:00
defaults.py Reduce and tighten code comments and docstrings repo-wide (#6095) 2026-06-08 23:09:51 -07:00
external_provider.py Studio: ignore unsupported env proxy during Studio startup (#6102) 2026-06-11 05:13:27 -07:00
inference.py Expose runtime context length for hub models (#6154) 2026-06-11 22:13:53 +03:00
key_exchange.py Reduce and tighten code comments and docstrings repo-wide (#6095) 2026-06-08 23:09:51 -07:00
llama_cpp.py Runtime MTP fallback for tensor parallelism (try MTP, recover if it crashes) (#6324) 2026-06-18 05:42:21 -07:00
llama_http.py Studio: serialize non-streaming responses once and pool the proxy client (#6393) 2026-06-17 22:38:02 -07:00
llama_server_args.py studio: deterministic VRAM auto-fit for GGUF (MTP reserve, compute buffer, total-based budget) (#6312) 2026-06-17 03:10:22 -07:00
llama_stats.py Studio: trim serving-log noise and surface llama-server engine stats (#6377) 2026-06-17 05:37:57 -07:00
mcp_client.py Studio: enable stdio MCP servers on a loopback bind (#6295) 2026-06-15 03:02:32 +01:00
mcp_config_import.py studio: show MCP "Import config" on the add-server form (#6030) 2026-06-11 16:17:22 +01:00
mlx_inference.py Expose runtime context length for hub models (#6154) 2026-06-11 22:13:53 +03:00
orchestrator.py Harden model fetching (#6391) 2026-06-18 05:39:52 -07:00
pricing.py Reduce and tighten code comments and docstrings repo-wide (#6095) 2026-06-08 23:09:51 -07:00
providers.py Studio: Add custom provider option to Connections (#6112) 2026-06-12 13:09:35 +02:00
runtime_context.py Expose runtime context length for hub models (#6154) 2026-06-11 22:13:53 +03:00
safetensors_agentic.py Studio: Bypass Permissions (skip confirmation, disable tool sandbox) (#5895) 2026-06-15 04:04:22 -07:00
tensor_fallback.py studio: deterministic VRAM auto-fit for GGUF (MTP reserve, compute buffer, total-based budget) (#6312) 2026-06-17 03:10:22 -07:00
tool_call_parser.py Studio: clean-room compact RAG (knowledge bases, hybrid search, fast indexing) (#5910) 2026-06-09 21:17:04 -07:00
tool_loop_controller.py Studio: clean-room compact RAG (knowledge bases, hybrid search, fast indexing) (#5910) 2026-06-09 21:17:04 -07:00
tools.py Rename chat artifacts copy to canvas (#6298) 2026-06-15 13:38:39 +02:00
worker.py Harden model fetching (#6391) 2026-06-18 05:39:52 -07:00