Merge remote-tracking branch 'origin/main' into studio/api-monitor-and-per-model-settings
# Conflicts: # studio/frontend/src/features/model-picker/components/sidebar-model-config.tsx
This commit is contained in:
commit
5d9d1b295d
102 changed files with 4034 additions and 244 deletions
9
.github/workflows/consolidated-tests-ci.yml
vendored
9
.github/workflows/consolidated-tests-ci.yml
vendored
|
|
@ -373,11 +373,10 @@ jobs:
|
|||
tests/test_bad_mappings_redirect.py \
|
||||
tests/test_prefetch_snapshot_scope.py \
|
||||
tests/test_gemma_2b_mapper_key.py \
|
||||
--deselect 'tests/utils/test_attention_masks.py::test_run_attention_flash_varlen_receives_window_and_softcap'
|
||||
# The deselected test monkeypatches flash_attn_varlen_func, which is
|
||||
# only bound on the module when `flash_attn` is importable. flash_attn
|
||||
# requires CUDA + dev toolchain, which the CPU-only ubuntu-latest
|
||||
# runner does not have. The other Bucket-A tests pass cleanly.
|
||||
tests/test_raw_text_json_loading.py
|
||||
# test_run_attention_flash_varlen_receives_window_and_softcap was deselected
|
||||
# until attention_dispatch.py predefined flash_attn_varlen_func as None; it
|
||||
# monkeypatches that name, so it no longer needs flash_attn on this runner.
|
||||
|
||||
- name: unsloth_zoo @ ${{ env.UNSLOTH_ZOO_REF }} — full pytest (CPU)
|
||||
# 106 of 111 test_* in unsloth_zoo are CPU-only. The two CUDA-skip
|
||||
|
|
|
|||
10
.github/workflows/studio-tauri-smoke.yml
vendored
10
.github/workflows/studio-tauri-smoke.yml
vendored
|
|
@ -91,6 +91,16 @@ jobs:
|
|||
npm run build
|
||||
test -f dist/index.html
|
||||
|
||||
# The crate carries ~100 unit tests (native_file_dialogs, preflight,
|
||||
# install, desktop_auth, ...) that nothing ran until now: this workflow
|
||||
# only ever built. Run them here, where the toolchain and the WebKit dev
|
||||
# packages are already installed, so a broken assertion fails the PR
|
||||
# instead of sitting unnoticed. `--no-fail-fast` reports every failing
|
||||
# test in one run rather than stopping at the first.
|
||||
- name: Rust unit tests (studio/src-tauri)
|
||||
working-directory: studio/src-tauri
|
||||
run: cargo test --no-fail-fast
|
||||
|
||||
- name: Tauri debug build (Linux, no bundle, no codesign)
|
||||
# `--debug` + `--no-bundle` keeps this lean: compiles the Rust crate,
|
||||
# confirms the frontend dist is wired into Tauri, but skips the AppImage
|
||||
|
|
|
|||
40
install.sh
40
install.sh
|
|
@ -321,10 +321,25 @@ _gfx906_bnb_prune() {
|
|||
|| "$_VENV_PY" -m pip uninstall -y bitsandbytes >/dev/null 2>&1 || true
|
||||
}
|
||||
|
||||
# Install bitsandbytes on AMD ROCm hosts. Uses the continuous-release_main
|
||||
# wheel for the ROCm 4-bit GEMV fix (bnb PR #1887, post-0.49.2); bnb <= 0.49.2
|
||||
# NaNs at decode shape on every AMD GPU. Falls back to PyPI >=0.49.1 if the
|
||||
# pre-release URL is unreachable. Drop the pin once bnb 0.50+ ships on PyPI.
|
||||
# Install bitsandbytes on AMD ROCm hosts. bnb <= 0.49.2 NaNs at 4-bit decode
|
||||
# shape on every AMD GPU; the fix (bnb #1887) ships in continuous-release_main
|
||||
# and, on PyPI, first in 0.50.0. Keep this floor in step with the amd extra in
|
||||
# pyproject.toml and studio/install_python_stack.py.
|
||||
_BNB_ROCM_PYPI_FALLBACK="bitsandbytes>=0.50.0"
|
||||
# bitsandbytes ships no ROCm binary in its aarch64 wheel at any version: the PyPI
|
||||
# 0.50.0 and continuous-release_main aarch64 wheels both carry only
|
||||
# libbitsandbytes_cpu.so plus CUDA variants. So neither install path below gives
|
||||
# aarch64 a 4-bit backend, and the messages must not claim one. Cf. gfx906.
|
||||
_bnb_rocm_arch_has_binary() {
|
||||
case "$_ARCH" in
|
||||
aarch64|arm64) return 1 ;;
|
||||
*) return 0 ;;
|
||||
esac
|
||||
}
|
||||
_warn_bnb_no_rocm_binary() {
|
||||
_bnb_rocm_arch_has_binary && return 0
|
||||
substep "[WARN] aarch64: bitsandbytes ships no ROCm kernels on this arch; 4-bit QLoRA needs a source build -- https://docs.unsloth.ai/get-started/install-and-update/amd" "$C_WARN"
|
||||
}
|
||||
_install_bnb_rocm() {
|
||||
_label="$1"
|
||||
_venv_py="$2"
|
||||
|
|
@ -339,9 +354,8 @@ _install_bnb_rocm() {
|
|||
_bnb_whl_url=""
|
||||
;;
|
||||
esac
|
||||
# uv rejects the continuous-release_main bitsandbytes wheel because the
|
||||
# filename version (1.33.7rc0) does not match the embedded metadata version
|
||||
# (0.50.0.dev0). pip accepts the mismatch, so bootstrap pip and use it.
|
||||
# uv rejects the pre-release wheel: filename version (1.33.7rc0) does not
|
||||
# match metadata (0.50.x.dev0). pip accepts it, so bootstrap pip and use it.
|
||||
if ! "$_venv_py" -m pip --version >/dev/null 2>&1; then
|
||||
if ! run_maybe_quiet "$_venv_py" -m ensurepip --upgrade; then
|
||||
run_maybe_quiet uv pip install --python "$_venv_py" pip || \
|
||||
|
|
@ -357,6 +371,7 @@ _install_bnb_rocm() {
|
|||
--retries 8 --timeout 90 \
|
||||
"$_bnb_whl_url" >"$_bnb_log" 2>&1; then
|
||||
rm -f "$_bnb_log"
|
||||
_warn_bnb_no_rocm_binary
|
||||
return 0
|
||||
fi
|
||||
_bnb_rc=$?
|
||||
|
|
@ -365,10 +380,17 @@ _install_bnb_rocm() {
|
|||
fi
|
||||
rm -f "$_bnb_log"
|
||||
step "warning" "$_label (pre-release) failed (exit code $_bnb_rc)" "$C_WARN" >&2
|
||||
substep "[WARN] bnb pre-release install failed; falling back to PyPI (4-bit decode broken on ROCm)" "$C_WARN"
|
||||
if _bnb_rocm_arch_has_binary; then
|
||||
substep "[WARN] bnb pre-release install failed; falling back to PyPI $_BNB_ROCM_PYPI_FALLBACK, which carries the ROCm 4-bit fix" "$C_WARN"
|
||||
else
|
||||
substep "[WARN] bnb pre-release install failed; falling back to PyPI $_BNB_ROCM_PYPI_FALLBACK" "$C_WARN"
|
||||
fi
|
||||
fi
|
||||
run_install_cmd "$_label (pypi fallback)" "$_venv_py" -m pip install \
|
||||
--force-reinstall --no-cache-dir --no-deps "bitsandbytes>=0.49.1"
|
||||
--force-reinstall --no-cache-dir --no-deps "$_BNB_ROCM_PYPI_FALLBACK"
|
||||
_bnb_pypi_rc=$?
|
||||
_warn_bnb_no_rocm_binary
|
||||
return $_bnb_pypi_rc
|
||||
}
|
||||
|
||||
if [ "$_next_is_package" = true ]; then
|
||||
|
|
|
|||
|
|
@ -1257,8 +1257,11 @@ intel = [
|
|||
]
|
||||
amd = [
|
||||
"unsloth[huggingfacenotorch]",
|
||||
"bitsandbytes>=0.49.1 ; ('linux' in sys_platform) and (platform_machine == 'AMD64' or platform_machine == 'x86_64' or platform_machine == 'aarch64')",
|
||||
"bitsandbytes>=0.49.1 ; (sys_platform == 'win32') and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
|
||||
# 4-bit decode is unreliable on ROCm before 0.50.0, the first PyPI release
|
||||
# carrying the full path: blocksize/warp decoupling (bnb #1887), fused SIMT
|
||||
# GEMM on RDNA (#1979), RDNA3/4 workgroup fix (#2012).
|
||||
"bitsandbytes>=0.50.0 ; ('linux' in sys_platform) and (platform_machine == 'AMD64' or platform_machine == 'x86_64' or platform_machine == 'aarch64')",
|
||||
"bitsandbytes>=0.50.0 ; (sys_platform == 'win32') and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
|
||||
]
|
||||
rocm702-torch280 = [
|
||||
"unsloth[amd]",
|
||||
|
|
|
|||
|
|
@ -58,6 +58,80 @@ DEFAULT_ADMISSION_QUEUE_PER_SLOT = 16
|
|||
DEFAULT_ADMISSION_MIN_QUEUE = 64
|
||||
|
||||
|
||||
def _executor_workers() -> int:
|
||||
"""Threads asyncio's default executor runs to_thread work on.
|
||||
|
||||
Mirrors ThreadPoolExecutor's own default sizing, which is what
|
||||
``run_in_executor(None, ...)`` builds. 3.13 sizes it from
|
||||
``process_cpu_count()``, which honours CPU affinity and cgroup quotas;
|
||||
``cpu_count()`` would budget from the whole host inside a one-core container.
|
||||
"""
|
||||
cpus = getattr(os, "process_cpu_count", os.cpu_count)() or 1
|
||||
return min(32, cpus + 4)
|
||||
|
||||
|
||||
def _executor_reserve(workers: int) -> int:
|
||||
"""Threads kept clear of parked approvals, for generation steps, stream
|
||||
teardown and unrelated to_thread work. Scaled rather than flat: a flat count
|
||||
would leave a 5-worker executor (one usable CPU) no budget at all.
|
||||
"""
|
||||
return max(2, workers // 8)
|
||||
|
||||
|
||||
def _max_parked(capacity: int) -> int:
|
||||
"""How many holders may sit on an approval prompt with their slot given back.
|
||||
|
||||
A pending prompt parks an executor thread (the loop blocks inside
|
||||
to_thread(next, gen)) whether or not it parked its slot, the pool already
|
||||
permits `capacity` of those, and every park admits one more, so budget only
|
||||
what the executor has left over. Zero on a backend whose --parallel alone
|
||||
fills it: the prompt then holds its slot, as it did before parking existed.
|
||||
"""
|
||||
workers = _executor_workers()
|
||||
spare = workers - _executor_reserve(workers) - max(0, capacity)
|
||||
# A quarter of the executor, floored at two while `spare` allows: a quarter of
|
||||
# five is one, and one park cannot cover the two simultaneous prompts #7455
|
||||
# exists for.
|
||||
return max(0, min(max(2, workers // 4), spare))
|
||||
|
||||
|
||||
# Process-wide, not per queue: there is one executor, and base_url takes a fresh
|
||||
# port on every load, so a per-queue budget would hand the same allowance to each
|
||||
# backend and to every reload, blind to the approvals parked on the old queue.
|
||||
_PARK_LOCK = threading.Lock()
|
||||
_parked_total = 0
|
||||
|
||||
|
||||
def _claim_park(limit: int) -> bool:
|
||||
global _parked_total
|
||||
with _PARK_LOCK:
|
||||
if _parked_total >= limit:
|
||||
return False
|
||||
_parked_total += 1
|
||||
return True
|
||||
|
||||
|
||||
def _drop_park() -> None:
|
||||
global _parked_total
|
||||
with _PARK_LOCK:
|
||||
_parked_total = max(0, _parked_total - 1)
|
||||
|
||||
|
||||
def _live_capacity(current: "LlamaAdmissionQueue") -> int:
|
||||
"""Slots across every backend still serving requests.
|
||||
|
||||
One queue's capacity is the wrong denominator for a budget sized against the
|
||||
one executor: a reload drains the old queue alongside the new one, and
|
||||
prompts on both park threads. Idle queues hold nothing and are about to be
|
||||
evicted.
|
||||
"""
|
||||
with _QUEUES_LOCK:
|
||||
queues = list(_QUEUES.values())
|
||||
# is_idle takes each queue's own lock, so never while holding _QUEUES_LOCK.
|
||||
total = sum(queue._capacity for queue in queues if queue is current or not queue.is_idle())
|
||||
return total if any(queue is current for queue in queues) else total + current._capacity
|
||||
|
||||
|
||||
@dataclass(frozen = True, **_SLOTS)
|
||||
class LlamaAdmissionConfig:
|
||||
enabled: bool = DEFAULT_ADMISSION_ENABLED
|
||||
|
|
@ -214,7 +288,7 @@ class _Waiter:
|
|||
|
||||
|
||||
class LlamaAdmissionLease:
|
||||
__slots__ = ("_queue", "_slot", "_released", "_release_lock", "_parked")
|
||||
__slots__ = ("_queue", "_slot", "_released", "_release_lock", "_parked", "_budgeted")
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -226,27 +300,52 @@ class LlamaAdmissionLease:
|
|||
self._released = False
|
||||
self._release_lock = threading.Lock()
|
||||
self._parked = False
|
||||
self._budgeted = False
|
||||
|
||||
@property
|
||||
def slot(self) -> Optional[int]:
|
||||
"""Pool slot this lease holds, or None when admission is disabled."""
|
||||
return self._slot
|
||||
|
||||
def park(self) -> None:
|
||||
def park(self) -> bool:
|
||||
"""Hand the slot back while this holder waits on something off the GPU.
|
||||
|
||||
A run stopped on a tool approval prompt is not decoding, so holding its
|
||||
slot would let unanswered prompts fill the pool while llama-server idles.
|
||||
The lease itself stays valid: releasing it after a park is still correct.
|
||||
|
||||
False when the park budget is spent and nothing was given back: the
|
||||
caller keeps its slot across the prompt, as it did before parking
|
||||
existed. Slower for whoever is behind it, but each freed slot admits
|
||||
another run that can park too, on the executor the generators run on.
|
||||
"""
|
||||
queue = self._queue
|
||||
slot = None
|
||||
with self._release_lock:
|
||||
if queue is None or self._released or self._parked:
|
||||
return
|
||||
return False
|
||||
# Under the lease lock so the decision and the handover cannot split.
|
||||
# Nothing takes the queue lock then a lease lock, so this order is
|
||||
# the only one in play.
|
||||
if not queue.try_park(self._slot):
|
||||
return False
|
||||
self._parked = True
|
||||
slot, self._slot = self._slot, None
|
||||
queue.park(slot)
|
||||
self._budgeted = True
|
||||
self._slot = None
|
||||
return True
|
||||
|
||||
def _drop_budget(self) -> None:
|
||||
"""Give the executor budget back now the prompt wait is over.
|
||||
|
||||
Separate from the queue's parked count, which lasts until the slot is
|
||||
back: the executor thread is free the moment the answer arrives. Holding
|
||||
the budget until the resume lands would refuse someone else's park for a
|
||||
finished wait, and that someone holds the slot the resumer wants.
|
||||
"""
|
||||
with self._release_lock:
|
||||
if not self._budgeted:
|
||||
return
|
||||
self._budgeted = False
|
||||
_drop_park()
|
||||
|
||||
def unpark(self) -> None:
|
||||
"""Drop the parked state without reclaiming a slot.
|
||||
|
|
@ -259,6 +358,7 @@ class LlamaAdmissionLease:
|
|||
if not self._parked:
|
||||
return
|
||||
self._parked = False
|
||||
self._drop_budget()
|
||||
if self._queue is not None:
|
||||
self._queue.unpark()
|
||||
|
||||
|
|
@ -278,6 +378,9 @@ class LlamaAdmissionLease:
|
|||
queue = self._queue
|
||||
if queue is None or not self._parked:
|
||||
return
|
||||
# Before the wait, not after: the prompt is answered, so this holder is
|
||||
# already off the executor and must not keep anyone else off it.
|
||||
self._drop_budget()
|
||||
slot = await queue.acquire_parked_slot(cancel_event = cancel_event, poll_s = poll_s)
|
||||
stranded = None
|
||||
with self._release_lock:
|
||||
|
|
@ -304,6 +407,7 @@ class LlamaAdmissionLease:
|
|||
self._released = True
|
||||
queue = self._queue
|
||||
parked, self._parked = self._parked, False
|
||||
self._drop_budget()
|
||||
if queue is not None:
|
||||
if parked:
|
||||
queue.unpark()
|
||||
|
|
@ -513,12 +617,20 @@ class LlamaAdmissionQueue:
|
|||
self._release_slot_locked(slot)
|
||||
self._grant_waiters_locked()
|
||||
|
||||
def park(self, slot: Optional[int]) -> None:
|
||||
"""Return a parked holder's slot to the pool. See ``LlamaAdmissionLease.park``."""
|
||||
def try_park(self, slot: Optional[int]) -> bool:
|
||||
"""Return a parked holder's slot to the pool. See ``LlamaAdmissionLease.park``.
|
||||
|
||||
False leaves the slot with its holder, so a refused park costs nothing to
|
||||
undo. The per-queue count is only what ``is_idle`` reads; the budget and
|
||||
the capacity it is sized from are both process-wide.
|
||||
"""
|
||||
if not _claim_park(_max_parked(_live_capacity(self))):
|
||||
return False
|
||||
with self._lock:
|
||||
self._parked += 1
|
||||
self._release_slot_locked(slot)
|
||||
self._grant_waiters_locked()
|
||||
return True
|
||||
|
||||
def unpark(self) -> None:
|
||||
with self._lock:
|
||||
|
|
@ -684,5 +796,10 @@ def get_llama_admission_queue(key: str) -> LlamaAdmissionQueue:
|
|||
|
||||
|
||||
def reset_llama_admission_queues() -> None:
|
||||
global _parked_total
|
||||
with _QUEUES_LOCK:
|
||||
_QUEUES.clear()
|
||||
# The budget outlives the queues it was claimed against, so dropping them
|
||||
# without it leaks the count and shrinks the budget for good.
|
||||
with _PARK_LOCK:
|
||||
_parked_total = 0
|
||||
|
|
|
|||
|
|
@ -92,6 +92,7 @@ from utils.subprocess_compat import (
|
|||
from utils.process_lifetime import child_popen_kwargs as _child_popen_kwargs
|
||||
from core.inference.tool_call_parser import (
|
||||
MAX_ACT_REPROMPTS as _MAX_REPROMPTS,
|
||||
NUDGE_TOOL_CALLS_STATUS as _NUDGE_TOOL_CALLS_STATUS,
|
||||
REPROMPT_MAX_CHARS as _REPROMPT_MAX_CHARS,
|
||||
is_short_intent_without_action as _is_short_intent_without_action,
|
||||
reprompt_to_act_message as _reprompt_to_act_message,
|
||||
|
|
@ -309,6 +310,15 @@ def _native_linux_system_rocm_lib_dirs(binary_dir: str = "") -> "list[str]":
|
|||
os.path.join(d, "libhsa-runtime64.so.1")
|
||||
):
|
||||
out.append(d)
|
||||
# ROCm keeps LLVM's versioned runtime under <root>/lib/llvm, so a
|
||||
# lib64 host still finds it under lib. Probe both and keep them
|
||||
# ahead of the bundle, else system libamd_comgr binds to the
|
||||
# bundle's incompatible libLLVM.so.*.
|
||||
for _sub in (lib_sub, "lib"):
|
||||
llvm_lib = os.path.join(base, _sub, "llvm", "lib")
|
||||
if llvm_lib not in seen and os.path.isdir(llvm_lib):
|
||||
seen.add(llvm_lib)
|
||||
out.append(llvm_lib)
|
||||
return out
|
||||
|
||||
|
||||
|
|
@ -1688,9 +1698,20 @@ def _kv_unified_from_args(
|
|||
return enabled
|
||||
|
||||
|
||||
def _flash_attn_enabled_from_args(args: Optional[Iterable[str]], default: bool = True) -> bool:
|
||||
"""Resolve llama.cpp's last-wins flash-attention CLI setting."""
|
||||
def _flash_attn_enabled_from_args(
|
||||
args: Optional[Iterable[str]],
|
||||
default: bool = True,
|
||||
env: Optional[Mapping[str, str]] = None,
|
||||
) -> bool:
|
||||
"""Resolve llama.cpp's environment and last-wins flash-attention settings."""
|
||||
enabled = default
|
||||
# llama.cpp applies LLAMA_ARG_FLASH_ATTN before parsing argv (arg.cpp set_env),
|
||||
# so the CLI still wins. --flash-attn has no args_neg, so no LLAMA_ARG_NO_ twin.
|
||||
value = (os.environ if env is None else env).get("LLAMA_ARG_FLASH_ATTN")
|
||||
if value in _LLAMA_ARG_FALSE_VALUES:
|
||||
enabled = False
|
||||
elif value in _LLAMA_ARG_TRUE_OR_AUTO_VALUES:
|
||||
enabled = True
|
||||
values = [str(arg) for arg in args] if args else []
|
||||
for i, raw in enumerate(values):
|
||||
if _flag_name(raw) not in {"-fa", "--flash-attn"}:
|
||||
|
|
@ -2172,6 +2193,8 @@ class LlamaCppBackend:
|
|||
self._effective_context_length: Optional[int] = None
|
||||
self._max_context_length: Optional[int] = None
|
||||
self._effective_parallel_slots: int = 1
|
||||
# --parallel the last load asked for, before any fit-time reduction.
|
||||
self._requested_n_parallel: int = 1
|
||||
self._chat_template: Optional[str] = None
|
||||
self._chat_template_override: Optional[str] = None
|
||||
self._supports_reasoning: bool = False
|
||||
|
|
@ -2408,6 +2431,17 @@ class LlamaCppBackend:
|
|||
slots = 1
|
||||
return max(1, slots)
|
||||
|
||||
@property
|
||||
def requested_parallel_slots(self) -> int:
|
||||
"""--parallel the last load asked for, before any fit-time reduction.
|
||||
The reload dedupe compares requested-vs-requested (like requested_n_ctx);
|
||||
the effective count would reload forever after a fitter reduction."""
|
||||
try:
|
||||
slots = int(getattr(self, "_requested_n_parallel", 1))
|
||||
except (TypeError, ValueError):
|
||||
slots = 1
|
||||
return max(1, slots)
|
||||
|
||||
@property
|
||||
def max_context_length(self) -> Optional[int]:
|
||||
"""Return the largest context that fits on this hardware at load time.
|
||||
|
|
@ -2433,6 +2467,8 @@ class LlamaCppBackend:
|
|||
|
||||
def _reset_effective_parallel_slots(self) -> None:
|
||||
self._effective_parallel_slots = 1
|
||||
# Cleared with the effective count so a stale value can't skew the dedupe.
|
||||
self._requested_n_parallel = 1
|
||||
|
||||
@staticmethod
|
||||
def _read_rss_bytes(pid: int) -> Optional[int]:
|
||||
|
|
@ -6778,6 +6814,7 @@ class LlamaCppBackend:
|
|||
chat_template_override = chat_template_override,
|
||||
extra_args = extra_args,
|
||||
is_vision = is_vision,
|
||||
n_parallel = n_parallel,
|
||||
preserve_multi_gpu_on_layer = preserve_multi_gpu_on_layer,
|
||||
):
|
||||
logger.info(
|
||||
|
|
@ -9029,7 +9066,8 @@ class LlamaCppBackend:
|
|||
int(self._DEFAULT_N_UBATCH if _effective_ubatch is None else _effective_ubatch),
|
||||
)
|
||||
self._flash_attn_enabled = (
|
||||
_flash_attn_enabled_from_args(_last_spawn_cmd) and self._architecture != "grok"
|
||||
_flash_attn_enabled_from_args(_last_spawn_cmd, env = env)
|
||||
and self._architecture != "grok"
|
||||
)
|
||||
self._effective_cache_types = _effective_main_cache_types(
|
||||
_last_spawn_cmd,
|
||||
|
|
@ -9057,6 +9095,8 @@ class LlamaCppBackend:
|
|||
self._extra_args = list(extra_args)
|
||||
self._extra_args_source = (model_identifier, hf_variant)
|
||||
self._requested_n_ctx = int(n_ctx)
|
||||
# Local n_parallel may have been reduced above; the snapshot has the ask.
|
||||
self._requested_n_parallel = max(1, int(_pending_load_kwargs["n_parallel"]))
|
||||
# Commit the known-good snapshot + whether MTP+tensor is live, then
|
||||
# watch this load for a mid-generation crash.
|
||||
self._last_load_kwargs = _pending_load_kwargs
|
||||
|
|
@ -9469,6 +9509,7 @@ class LlamaCppBackend:
|
|||
tensor_split: Optional[List[float]] = None,
|
||||
gpu_ids: Optional[List[int]] = None,
|
||||
mtp_draft_path: Optional[str] = None,
|
||||
n_parallel: int = 1,
|
||||
preserve_multi_gpu_on_layer: bool = False,
|
||||
) -> bool:
|
||||
"""True iff the live server already satisfies these load kwargs.
|
||||
|
|
@ -9533,6 +9574,10 @@ class LlamaCppBackend:
|
|||
# A GPU-memory-mode flip (Unsloth / manual) must always reload.
|
||||
if self._gpu_memory_mode != gpu_memory_mode:
|
||||
return False
|
||||
# Requested-vs-requested (like n_ctx): comparing the effective count
|
||||
# would reload forever whenever the fitter launched fewer slots.
|
||||
if self._requested_n_parallel != max(1, int(n_parallel)):
|
||||
return False
|
||||
# Manual: a layer-count change always reloads (covers Auto(-1) <-> a
|
||||
# pinned count); MoE/split only matter with an explicit offload.
|
||||
if gpu_memory_mode == "manual" and (
|
||||
|
|
@ -12375,7 +12420,10 @@ class LlamaCppBackend:
|
|||
_it_r = _iter_timings or {}
|
||||
_accumulated_predicted_ms += _it_r.get("predicted_ms", 0)
|
||||
_accumulated_predicted_n += _it_r.get("predicted_n", 0)
|
||||
# Blank first (the route resets its text cursor only on an
|
||||
# empty status), then the badge so the retry is not a hang.
|
||||
yield {"type": "status", "text": ""}
|
||||
yield {"type": "status", "text": _NUDGE_TOOL_CALLS_STATUS}
|
||||
continue
|
||||
|
||||
if _forced_tool_call_pending:
|
||||
|
|
|
|||
|
|
@ -16,11 +16,18 @@ from __future__ import annotations
|
|||
import os
|
||||
from typing import Iterable, Mapping, Optional
|
||||
|
||||
# Valid llama-server --parallel range, shared with LoadRequest.n_parallel.
|
||||
# Mirrored by callers that cannot import this: run.py and unsloth_cli/commands/
|
||||
# studio.py (_PARALLEL_MIN/MAX), per-model-config.ts (N_PARALLEL_MIN/MAX);
|
||||
# test_parallel_slots_per_load.py pins them together.
|
||||
PARALLEL_MIN = 1
|
||||
PARALLEL_MAX = 64
|
||||
|
||||
# Each group = every alias (short + long) of one hard-denied flag.
|
||||
# Extend the matching group when llama.cpp adds a new alias.
|
||||
_DENYLIST_GROUPS: tuple[frozenset[str], ...] = (
|
||||
# Parallel slots: owned by typer --parallel; a pass-through would desync
|
||||
# app.state.llama_parallel_slots from llama-server.
|
||||
# Parallel slots: owned by typer --parallel and LoadRequest.n_parallel; a
|
||||
# pass-through would desync the slot bookkeeping from llama-server.
|
||||
frozenset({"-np", "--parallel", "--n-parallel"}),
|
||||
# Model identity: Unsloth resolves it from LoadRequest; a second -m would
|
||||
# load a different model than Unsloth thinks it loaded.
|
||||
|
|
|
|||
|
|
@ -971,7 +971,12 @@ def _call_stdio_tool(
|
|||
raise RuntimeError("MCP server connection is not available")
|
||||
else:
|
||||
rem = _remaining()
|
||||
coro = _race_tool_call(session.client.call_tool(name, args), rem, cancel_event)
|
||||
# raise_on_error=False for the same reason as the one-shot path.
|
||||
coro = _race_tool_call(
|
||||
session.client.call_tool(name, args, raise_on_error = False),
|
||||
rem,
|
||||
cancel_event,
|
||||
)
|
||||
return session.run(coro, rem)
|
||||
except (_MCPCancelled, asyncio.TimeoutError):
|
||||
# _race_tool_call cancels the pending call but cancellation is
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ from core.inference.tool_call_parser import (
|
|||
_strip_mistral_reasoning,
|
||||
BUDGET_EXHAUSTED_NUDGE,
|
||||
MAX_ACT_REPROMPTS,
|
||||
NUDGE_TOOL_CALLS_STATUS,
|
||||
RAG_MAX_SEARCHES_PER_TURN,
|
||||
RAG_SEARCH_CAP_NUDGE,
|
||||
TOOL_XML_SIGNALS,
|
||||
|
|
@ -1032,9 +1033,10 @@ def run_safetensors_tool_loop(
|
|||
"content": reprompt_to_act_message(tool_hint),
|
||||
}
|
||||
)
|
||||
# Empty status clears the badge and resets the route's
|
||||
# per-turn text cursor before the re-prompted turn streams.
|
||||
# Blank first: it clears the badge and resets the route's per-turn
|
||||
# text cursor. The badge then shows the pause is a re-prompt, not a stall.
|
||||
yield {"type": "status", "text": ""}
|
||||
yield {"type": "status", "text": NUDGE_TOOL_CALLS_STATUS}
|
||||
continue
|
||||
|
||||
# Final answer. If a literal tool marker in prose was buffered but
|
||||
|
|
|
|||
|
|
@ -183,6 +183,9 @@ INTENT_SIGNAL = re.compile(
|
|||
# times since #5620); safetensors and MLX inherit the same cap from here.
|
||||
MAX_ACT_REPROMPTS = 3
|
||||
REPROMPT_MAX_CHARS = 2000
|
||||
# Composer badge while a hidden re-prompted turn regenerates, else the UI looks
|
||||
# hung. Matched exactly by the frontend (utils/tool-status.ts); keep in sync.
|
||||
NUDGE_TOOL_CALLS_STATUS = "Nudging tool calls"
|
||||
|
||||
|
||||
def is_short_intent_without_action(text: str) -> bool:
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from pydantic import (
|
|||
model_validator,
|
||||
)
|
||||
|
||||
from core.inference.llama_server_args import PARALLEL_MAX, PARALLEL_MIN
|
||||
from picker.schemas import MAX_CHAT_TEMPLATE_BYTES
|
||||
|
||||
|
||||
|
|
@ -113,6 +114,18 @@ class LoadRequest(BaseModel):
|
|||
"'mtp' or 'mtp+ngram'."
|
||||
),
|
||||
)
|
||||
n_parallel: Optional[int] = Field(
|
||||
None,
|
||||
ge = PARALLEL_MIN,
|
||||
le = PARALLEL_MAX,
|
||||
description = (
|
||||
"Parallel decode slots for llama-server (--parallel) for this "
|
||||
f"load ({PARALLEL_MIN}..{PARALLEL_MAX}). Omit for the server-wide "
|
||||
"default set at launch (the --parallel CLI flag). The VRAM fitter "
|
||||
"may launch fewer slots to keep the model fully on GPU. Ignored "
|
||||
"for non-GGUF models."
|
||||
),
|
||||
)
|
||||
tensor_parallel: bool = Field(
|
||||
False,
|
||||
description = (
|
||||
|
|
@ -265,6 +278,16 @@ class ValidateModelRequest(BaseModel):
|
|||
"delegate fitting to llama.cpp, while explicit layers are user-owned."
|
||||
),
|
||||
)
|
||||
n_parallel: Optional[int] = Field(
|
||||
None,
|
||||
ge = PARALLEL_MIN,
|
||||
le = PARALLEL_MAX,
|
||||
description = (
|
||||
"Parallel decode slots intended for the follow-up load, so the "
|
||||
"coexistence estimate sizes the KV cache like /load. Omit for the "
|
||||
"server-wide --parallel default."
|
||||
),
|
||||
)
|
||||
include_context_length: bool = Field(
|
||||
False,
|
||||
description = "Also read the native context length from the local GGUF header. "
|
||||
|
|
@ -533,6 +556,23 @@ class LoadResponse(BaseModel):
|
|||
"or None for automatic selection."
|
||||
),
|
||||
)
|
||||
requested_parallel_slots: Optional[int] = Field(
|
||||
None,
|
||||
description = (
|
||||
"Parallel decode slots the load was invoked with (per-load "
|
||||
"n_parallel, else the server-wide --parallel default). None for "
|
||||
"non-GGUF loads and for the diffusion runner, which ignores "
|
||||
"--parallel."
|
||||
),
|
||||
)
|
||||
parallel_slots: Optional[int] = Field(
|
||||
None,
|
||||
description = (
|
||||
"Serving slots the active llama-server actually runs (--parallel "
|
||||
"after any fit-time slot reduction). None for non-GGUF loads and "
|
||||
"for the diffusion runner, which ignores --parallel."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class UnloadResponse(BaseModel):
|
||||
|
|
@ -708,6 +748,23 @@ class InferenceStatusResponse(BaseModel):
|
|||
"or None for automatic selection."
|
||||
),
|
||||
)
|
||||
requested_parallel_slots: Optional[int] = Field(
|
||||
None,
|
||||
description = (
|
||||
"Parallel decode slots the active load was invoked with (per-load "
|
||||
"n_parallel, else the server-wide --parallel default). None when "
|
||||
"no GGUF model is loaded and for the diffusion runner, which "
|
||||
"ignores --parallel."
|
||||
),
|
||||
)
|
||||
parallel_slots: Optional[int] = Field(
|
||||
None,
|
||||
description = (
|
||||
"Serving slots the active llama-server actually runs (--parallel "
|
||||
"after any fit-time slot reduction). None when no GGUF model is "
|
||||
"loaded and for the diffusion runner, which ignores --parallel."
|
||||
),
|
||||
)
|
||||
llama_cpp_supports_mtp: bool = Field(
|
||||
True,
|
||||
description = (
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
|||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
|
||||
from auth.authentication import get_current_subject
|
||||
from core.inference.llama_server_args import PARALLEL_MAX, PARALLEL_MIN
|
||||
from loggers import get_logger
|
||||
from utils.utils import safe_curated_detail, log_and_http_error
|
||||
from storage.studio_db import (
|
||||
|
|
@ -169,6 +170,7 @@ class ChatPresetLoadConfig(BaseModel):
|
|||
kvCacheDtype: Optional[str] = None
|
||||
speculativeType: Optional[str] = None
|
||||
specDraftNMax: Optional[int] = Field(default = None, ge = 1, le = 16)
|
||||
nParallel: Optional[int] = Field(default = None, ge = PARALLEL_MIN, le = PARALLEL_MAX)
|
||||
tensorParallel: Optional[bool] = None
|
||||
gpuMemoryMode: Optional[Literal["manual"]] = None
|
||||
gpuLayers: Optional[int] = None
|
||||
|
|
|
|||
|
|
@ -3311,10 +3311,25 @@ def _is_explicit_tensor_drop(request: LoadRequest) -> bool:
|
|||
return override is not None and override.strip().lower() != "tensor"
|
||||
|
||||
|
||||
def _parallel_slot_echo(llama_backend: LlamaCppBackend) -> dict:
|
||||
"""requested/effective parallel-slot fields for /load and /status echoes.
|
||||
|
||||
The diffusion runner ignores ``--parallel`` and never commits a count, so it
|
||||
reports None like the non-GGUF paths; echoing the reset placeholder 1 would
|
||||
fabricate an "invoked with 1 slot"."""
|
||||
if llama_backend.is_diffusion:
|
||||
return {"requested_parallel_slots": None, "parallel_slots": None}
|
||||
return {
|
||||
"requested_parallel_slots": llama_backend.requested_parallel_slots,
|
||||
"parallel_slots": llama_backend.effective_parallel_slots,
|
||||
}
|
||||
|
||||
|
||||
def _request_matches_loaded_settings(
|
||||
request: LoadRequest,
|
||||
llama_backend: LlamaCppBackend,
|
||||
effective_chat_template_override: Optional[str] = None,
|
||||
requested_parallel_slots: Optional[int] = None,
|
||||
) -> bool:
|
||||
"""True iff every runtime setting on the request matches the loaded server.
|
||||
Caller has already checked model+variant+is_loaded. See #5401.
|
||||
|
|
@ -3323,11 +3338,22 @@ def _request_matches_loaded_settings(
|
|||
launched (user override, else a bundled family template such as the
|
||||
gemma-4 override), so the dedup compares against what the backend actually
|
||||
holds rather than the raw request field. Defaults to the request field for
|
||||
callers that do not resolve a bundled override."""
|
||||
callers that do not resolve a bundled override.
|
||||
|
||||
``requested_parallel_slots`` is the resolved count the load would use
|
||||
(per-load ``n_parallel``, else the server-wide default); None skips it."""
|
||||
# Compare requested n_ctx (not effective) so VRAM-cap doesn't mask an
|
||||
# Auto-vs-explicit slider flip.
|
||||
if request.max_seq_length != llama_backend.requested_n_ctx:
|
||||
return False
|
||||
# Requested-vs-requested for the same reason: the fitter may launch fewer
|
||||
# slots. Diffusion ignores --parallel, so a change there must not reload.
|
||||
if (
|
||||
requested_parallel_slots is not None
|
||||
and not llama_backend.is_diffusion
|
||||
and int(requested_parallel_slots) != llama_backend.requested_parallel_slots
|
||||
):
|
||||
return False
|
||||
if _normalise_settings_str(request.cache_type_kv) != _normalise_settings_str(
|
||||
llama_backend.cache_type_kv
|
||||
):
|
||||
|
|
@ -4865,6 +4891,20 @@ def _guard_chat_load_against_training(
|
|||
cpu_only = LlamaCppBackend._effective_gpu_count() == 0,
|
||||
)
|
||||
|
||||
# Size with the count that will actually launch, or a load that fits gets a
|
||||
# 409: diffusion never receives --parallel, and load_model clamps to 1 on an
|
||||
# llama-server without --kv-unified. An unclassified GGUF keeps the ask.
|
||||
if is_gguf and n_parallel > 1:
|
||||
if diffusion_kind is True:
|
||||
n_parallel = 1
|
||||
else:
|
||||
try:
|
||||
caps = LlamaCppBackend.probe_server_capabilities()
|
||||
if caps.get("found") and not caps.get("supports_kv_unified"):
|
||||
n_parallel = 1
|
||||
except Exception as e:
|
||||
logger.warning("Could not probe llama-server slots for chat-load guard: %s", e)
|
||||
|
||||
required_override_gb = (
|
||||
_estimate_gguf_required_gb(
|
||||
config,
|
||||
|
|
@ -5413,6 +5453,17 @@ async def _load_model_impl(
|
|||
backend = get_inference_backend()
|
||||
llama_backend = get_llama_cpp_backend()
|
||||
|
||||
# Resolve the slot count once (per-load field, else the server-wide
|
||||
# --parallel default) so the dedupe, the training guard and the load
|
||||
# kwargs all size against what launches. app.state stays the launch
|
||||
# intent / admission fallback; getattr because direct callers have no app.
|
||||
_app_state = getattr(getattr(fastapi_request, "app", None), "state", None)
|
||||
_n_parallel = (
|
||||
request.n_parallel
|
||||
if request.n_parallel is not None
|
||||
else getattr(_app_state, "llama_parallel_slots", 1)
|
||||
)
|
||||
|
||||
is_direct_gguf_request = model_identifier.lower().endswith(".gguf")
|
||||
if request.gguf_variant or is_direct_gguf_request:
|
||||
gguf_variant_matches = is_direct_gguf_request or bool(
|
||||
|
|
@ -5430,6 +5481,7 @@ async def _load_model_impl(
|
|||
request,
|
||||
llama_backend,
|
||||
effective_chat_template_override,
|
||||
requested_parallel_slots = _n_parallel,
|
||||
)
|
||||
# Skip if a prior audio probe failed -- let load_model retry.
|
||||
and getattr(llama_backend, "_audio_probed", True)
|
||||
|
|
@ -5484,6 +5536,7 @@ async def _load_model_impl(
|
|||
n_moe_layers = llama_backend.n_moe_layers,
|
||||
gpu_ids = llama_backend.gpu_ids,
|
||||
requested_gpu_ids = llama_backend.requested_gpu_ids,
|
||||
**_parallel_slot_echo(llama_backend),
|
||||
)
|
||||
else:
|
||||
if (
|
||||
|
|
@ -5622,7 +5675,7 @@ async def _load_model_impl(
|
|||
max_seq_length = request.max_seq_length,
|
||||
requested_gpu_ids = effective_gpu_ids,
|
||||
llama_extra_args = extra_llama_args,
|
||||
n_parallel = getattr(fastapi_request.app.state, "llama_parallel_slots", 1),
|
||||
n_parallel = _n_parallel,
|
||||
cache_type_kv = request.cache_type_kv,
|
||||
tensor_parallel = bool(request.tensor_parallel),
|
||||
gpu_memory_mode = request.gpu_memory_mode,
|
||||
|
|
@ -5699,7 +5752,6 @@ async def _load_model_impl(
|
|||
# Route to HF or local mode based on config. Run in a thread so the
|
||||
# event loop stays free for progress polling and other requests
|
||||
# during the (potentially long) GGUF download + llama-server start.
|
||||
_n_parallel = getattr(fastapi_request.app.state, "llama_parallel_slots", 1)
|
||||
|
||||
# Load kwargs common to HF and local modes; the two differ only by
|
||||
# the model-source args (hf_repo/-token vs gguf_path/mmproj).
|
||||
|
|
@ -5897,6 +5949,7 @@ async def _load_model_impl(
|
|||
n_moe_layers = llama_backend.n_moe_layers,
|
||||
gpu_ids = llama_backend.gpu_ids,
|
||||
requested_gpu_ids = llama_backend.requested_gpu_ids,
|
||||
**_parallel_slot_echo(llama_backend),
|
||||
)
|
||||
|
||||
# ── Standard path: load via Unsloth/transformers ──────────
|
||||
|
|
@ -6297,9 +6350,14 @@ async def validate_model(
|
|||
requested_gpu_ids = effective_gpu_ids,
|
||||
llama_extra_args = effective_extra_args,
|
||||
n_parallel = (
|
||||
getattr(fastapi_request.app.state, "llama_parallel_slots", 1)
|
||||
if fastapi_request is not None
|
||||
else 1
|
||||
request.n_parallel
|
||||
if request.n_parallel is not None
|
||||
# Same getattr chain as the load path: preflight must size like the load.
|
||||
else getattr(
|
||||
getattr(getattr(fastapi_request, "app", None), "state", None),
|
||||
"llama_parallel_slots",
|
||||
1,
|
||||
)
|
||||
),
|
||||
cache_type_kv = request.cache_type_kv,
|
||||
tensor_parallel = request.tensor_parallel,
|
||||
|
|
@ -7144,6 +7202,7 @@ async def get_status(current_subject: str = Depends(get_current_subject)):
|
|||
n_moe_layers = llama_backend.n_moe_layers,
|
||||
gpu_ids = llama_backend.gpu_ids,
|
||||
requested_gpu_ids = llama_backend.requested_gpu_ids,
|
||||
**_parallel_slot_echo(llama_backend),
|
||||
llama_cpp_supports_mtp = _supports_mtp,
|
||||
spec_fallback_reason = llama_backend.spec_fallback_reason,
|
||||
llama_cpp_prebuilt_stale = _stale,
|
||||
|
|
@ -9574,7 +9633,10 @@ async def openai_chat_completions(
|
|||
if lease is None:
|
||||
return
|
||||
if on:
|
||||
lease.park()
|
||||
# Refused when the budget is spent: the slot stays here,
|
||||
# so there is nothing to take back afterwards.
|
||||
if not lease.park():
|
||||
return
|
||||
elif wait:
|
||||
# Resuming: park() may have handed our slot to a waiter, so wait for room instead
|
||||
# of putting two holders on one slot.
|
||||
|
|
|
|||
|
|
@ -1920,7 +1920,8 @@ def _build_arg_parser():
|
|||
default = _PARALLEL_DEFAULT_PLAIN,
|
||||
help = (
|
||||
f"llama-server parallel decode slots ({_PARALLEL_MIN}..{_PARALLEL_MAX}). "
|
||||
f"Default {_PARALLEL_DEFAULT_PLAIN}."
|
||||
f"Default {_PARALLEL_DEFAULT_PLAIN}. The Studio run settings "
|
||||
"(Parallel Slots) override it per load."
|
||||
),
|
||||
)
|
||||
return parser
|
||||
|
|
|
|||
|
|
@ -1066,3 +1066,229 @@ def test_an_immediate_arrival_cannot_take_an_approved_chats_slot():
|
|||
assert queue.snapshot().active <= 1
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_parking_is_bounded_so_the_thread_pool_cannot_be_drained(monkeypatch):
|
||||
# A pending prompt parks an executor thread (the loop blocks inside
|
||||
# to_thread(next, gen)) and frees a slot that admits another run which can
|
||||
# park too, so unbounded parking drains the pool the generators run on.
|
||||
# Pinned because the real budget follows the runner's usable CPUs.
|
||||
monkeypatch.setattr(llama_admission, "_executor_workers", lambda: 32)
|
||||
|
||||
async def scenario():
|
||||
queue = get_llama_admission_queue("http://llama.test")
|
||||
config = LlamaAdmissionConfig()
|
||||
limit = llama_admission._max_parked(1)
|
||||
assert limit >= 1
|
||||
|
||||
leases = []
|
||||
for _ in range(limit):
|
||||
lease = queue.reserve(capacity = 1, config = config).lease_nowait()
|
||||
assert lease is not None and lease.park()
|
||||
leases.append(lease)
|
||||
|
||||
refused = queue.reserve(capacity = 1, config = config).lease_nowait()
|
||||
assert refused is not None
|
||||
assert not refused.park(), "parking is unbounded"
|
||||
# Refusing means keeping the slot, the old behaviour, not an error.
|
||||
assert refused.slot is not None
|
||||
assert queue.snapshot().active == 1
|
||||
|
||||
leases[0].unpark()
|
||||
assert refused.park(), "budget was not returned"
|
||||
for lease in leases[1:] + [refused]:
|
||||
lease.release()
|
||||
leases[0].release()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_the_park_budget_is_shared_by_every_queue(monkeypatch):
|
||||
# One executor, so a per-queue budget would be handed out again to every
|
||||
# backend and to every reload onto a fresh ephemeral port.
|
||||
monkeypatch.setattr(llama_admission, "_executor_workers", lambda: 32)
|
||||
|
||||
async def scenario():
|
||||
config = LlamaAdmissionConfig()
|
||||
first = get_llama_admission_queue("http://llama.test:1")
|
||||
second = get_llama_admission_queue("http://llama.test:2")
|
||||
limit = llama_admission._max_parked(1)
|
||||
|
||||
for index in range(limit):
|
||||
queue = first if index % 2 == 0 else second
|
||||
lease = queue.reserve(capacity = 1, config = config).lease_nowait()
|
||||
assert lease.park()
|
||||
|
||||
spare = second.reserve(capacity = 1, config = config).lease_nowait()
|
||||
assert not spare.park(), "each queue got its own budget"
|
||||
|
||||
# A reset drops the queues the count was claimed against, so it must drop
|
||||
# the count too or the leak shrinks the budget process-wide.
|
||||
reset_llama_admission_queues()
|
||||
revived = get_llama_admission_queue("http://llama.test:1")
|
||||
fresh = revived.reserve(capacity = 1, config = config).lease_nowait()
|
||||
assert fresh.park(), "reset leaked the park count"
|
||||
fresh.release()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_the_park_budget_leaves_the_executor_room_to_work(monkeypatch):
|
||||
# The pool already permits `capacity` pending prompts and every park admits
|
||||
# one more, so the budget must account for both. Swept across executor sizes
|
||||
# rather than read off this host, since a container gets a small one.
|
||||
for cpus in (1, 2, 4, 8, 16, 28, 64):
|
||||
workers = min(32, cpus + 4)
|
||||
monkeypatch.setattr(llama_admission, "_executor_workers", lambda w = workers: w)
|
||||
reserve = llama_admission._executor_reserve(workers)
|
||||
assert reserve >= 2, f"{workers} workers left no reserve"
|
||||
|
||||
# Even the smallest executor fits the two simultaneous prompts #7455 needs.
|
||||
assert llama_admission._max_parked(1) >= 2, f"no room for two on {workers} workers"
|
||||
assert llama_admission._max_parked(1) <= workers // 2
|
||||
# A backend whose --parallel alone fills the executor gets no parks.
|
||||
assert llama_admission._max_parked(workers) == 0
|
||||
for capacity in range(0, workers + 8):
|
||||
budget = llama_admission._max_parked(capacity)
|
||||
assert budget >= 0, f"negative budget at capacity {capacity}"
|
||||
assert (
|
||||
budget == 0 or capacity + budget <= workers - reserve
|
||||
), f"{workers} workers: capacity {capacity} plus {budget} parks leaves no room"
|
||||
|
||||
|
||||
def test_the_park_budget_follows_the_executors_own_cpu_count(monkeypatch):
|
||||
# 3.13 sizes ThreadPoolExecutor from process_cpu_count(), which honours CPU
|
||||
# affinity and cgroup quotas; cpu_count() would budget from the whole host
|
||||
# inside a one-core container. Pulled apart here, since they usually match.
|
||||
import concurrent.futures
|
||||
|
||||
monkeypatch.setattr(os, "cpu_count", lambda: 64)
|
||||
if hasattr(os, "process_cpu_count"):
|
||||
monkeypatch.setattr(os, "process_cpu_count", lambda: 1)
|
||||
# Against the real thing rather than the formula: the default executor is a
|
||||
# plain ThreadPoolExecutor(), so its own sizing is the answer on any version.
|
||||
with concurrent.futures.ThreadPoolExecutor() as pool:
|
||||
assert llama_admission._executor_workers() == pool._max_workers
|
||||
|
||||
|
||||
def test_the_stream_retries_a_park_that_was_refused():
|
||||
# _park_admission short-circuits on `on == _parked`, so recording a refused
|
||||
# park as parked would skip every later approval in the run even once the
|
||||
# budget frees up. Structural because that only shows on a second approval.
|
||||
import ast
|
||||
|
||||
# Read rather than import: routes.inference pulls in the whole app.
|
||||
route = os.path.join(_backend, "routes", "inference.py")
|
||||
with open(route, encoding = "utf-8") as handle:
|
||||
tree = ast.parse(handle.read())
|
||||
helpers = [
|
||||
node
|
||||
for node in ast.walk(tree)
|
||||
if isinstance(node, ast.AsyncFunctionDef) and node.name == "_park_admission"
|
||||
]
|
||||
assert len(helpers) == 1, f"expected one _park_admission, found {len(helpers)}"
|
||||
|
||||
guards = [
|
||||
node
|
||||
for node in ast.walk(helpers[0])
|
||||
if isinstance(node, ast.If)
|
||||
and isinstance(node.test, ast.UnaryOp)
|
||||
and isinstance(node.test.op, ast.Not)
|
||||
and isinstance(node.test.operand, ast.Call)
|
||||
and getattr(node.test.operand.func, "attr", None) == "park"
|
||||
and getattr(node.test.operand.func.value, "id", None) == "lease"
|
||||
]
|
||||
assert len(guards) == 1, "lease.park()'s answer is ignored"
|
||||
assert all(
|
||||
isinstance(stmt, ast.Return) for stmt in guards[0].body
|
||||
), "a refused park must leave _parked alone, so a later approval retries it"
|
||||
|
||||
|
||||
def test_the_park_budget_counts_every_live_backend(monkeypatch):
|
||||
# base_url takes a fresh port on every load, so a reload mints a queue while
|
||||
# the old one drains. Prompts on both park threads of the one executor, so a
|
||||
# budget sized from either backend alone lets them add up past the reserve.
|
||||
monkeypatch.setattr(llama_admission, "_executor_workers", lambda: 32)
|
||||
|
||||
async def scenario():
|
||||
config = LlamaAdmissionConfig()
|
||||
old = get_llama_admission_queue("http://llama.test:1")
|
||||
draining = old.reserve(capacity = 16, config = config).lease_nowait()
|
||||
assert draining is not None # in flight, so the registry keeps this queue
|
||||
|
||||
new = get_llama_admission_queue("http://llama.test:2")
|
||||
lease = new.reserve(capacity = 16, config = config).lease_nowait()
|
||||
assert lease is not None
|
||||
|
||||
# 16 slots each against 32 workers: their prompts alone can fill it.
|
||||
assert llama_admission._max_parked(16) > 0, "this test needs a budget to remove"
|
||||
assert not lease.park(), "budget sized from one backend of two"
|
||||
|
||||
draining.release() # the old backend drains and is up for eviction
|
||||
assert lease.park(), "an idle backend still counted against the budget"
|
||||
lease.release()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_the_park_budget_is_freed_when_the_prompt_is_answered(monkeypatch):
|
||||
# The executor thread comes back the moment the answer arrives, before the
|
||||
# resume queues for a slot. Holding the budget until the slot lands refuses
|
||||
# someone else's park, and that someone holds the slot the resumer wants.
|
||||
monkeypatch.setattr(llama_admission, "_executor_workers", lambda: 32)
|
||||
|
||||
async def scenario():
|
||||
config = LlamaAdmissionConfig()
|
||||
queue = get_llama_admission_queue("http://llama.test")
|
||||
|
||||
parked = []
|
||||
for _ in range(llama_admission._max_parked(1)):
|
||||
lease = queue.reserve(capacity = 1, config = config).lease_nowait()
|
||||
assert lease is not None and lease.park()
|
||||
parked.append(lease)
|
||||
|
||||
blocked = queue.reserve(capacity = 1, config = config).lease_nowait()
|
||||
assert blocked is not None
|
||||
assert not blocked.park(), "the budget was not full to begin with"
|
||||
|
||||
# One prompt is answered. Its slot is taken, so the resume queues for one.
|
||||
resumed = asyncio.ensure_future(parked[0].unpark_async(poll_s = 0.01))
|
||||
await asyncio.sleep(0.05)
|
||||
assert not resumed.done(), "the resume needs to still be waiting for its slot"
|
||||
|
||||
assert blocked.park(), "budget held for a prompt wait that is over"
|
||||
# Which is what frees the slot the resumer was waiting for.
|
||||
await asyncio.wait_for(resumed, timeout = 2)
|
||||
for lease in parked[1:] + [blocked]:
|
||||
lease.release()
|
||||
parked[0].release()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_releasing_a_parked_holder_returns_its_budget(monkeypatch):
|
||||
# A client that disconnects on the prompt releases straight out of parked,
|
||||
# never unparking. Its executor thread went with it, so keeping the budget
|
||||
# would lose one for the life of the process.
|
||||
monkeypatch.setattr(llama_admission, "_executor_workers", lambda: 32)
|
||||
|
||||
async def scenario():
|
||||
config = LlamaAdmissionConfig()
|
||||
queue = get_llama_admission_queue("http://llama.test")
|
||||
|
||||
parked = []
|
||||
for _ in range(llama_admission._max_parked(1)):
|
||||
lease = queue.reserve(capacity = 1, config = config).lease_nowait()
|
||||
assert lease is not None and lease.park()
|
||||
parked.append(lease)
|
||||
|
||||
blocked = queue.reserve(capacity = 1, config = config).lease_nowait()
|
||||
assert blocked is not None
|
||||
assert not blocked.park(), "the budget was not full to begin with"
|
||||
|
||||
parked[0].release()
|
||||
assert blocked.park(), "a released park never gave its budget back"
|
||||
for lease in parked[1:] + [blocked]:
|
||||
lease.release()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from core.inference.llama_cpp import (
|
|||
_PROVISIONAL_ARGS_MIN_CHARS,
|
||||
LlamaCppBackend,
|
||||
)
|
||||
from core.inference.tool_call_parser import NUDGE_TOOL_CALLS_STATUS
|
||||
from state import tool_approvals
|
||||
from state.tool_approvals import TOOL_REJECTED_MESSAGE, resolve_tool_decision
|
||||
|
||||
|
|
@ -1841,6 +1842,140 @@ def test_reprompted_tool_call_still_streams_final_answer(monkeypatch):
|
|||
assert len(payloads) == 3
|
||||
|
||||
|
||||
def _status_texts(events: list[dict]) -> list[str]:
|
||||
return [event["text"] for event in events if event.get("type") == "status"]
|
||||
|
||||
|
||||
_WEB_SEARCH_TOOL = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "web_search",
|
||||
"description": "Search the web.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"query": {"type": "string"}},
|
||||
"required": ["query"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _nudge_then_search_streams() -> list[list[str]]:
|
||||
"""Stall, then a re-prompted turn that finally searches, then the answer."""
|
||||
|
||||
return [
|
||||
[_sse({"content": "I will search the web now."}), _done()],
|
||||
[
|
||||
_sse(
|
||||
{
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": 0,
|
||||
"id": "call_search",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "web_search",
|
||||
"arguments": json.dumps({"query": "red square"}),
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
),
|
||||
_done(),
|
||||
],
|
||||
[_sse({"content": "Final answer: the square is red."}), _done()],
|
||||
]
|
||||
|
||||
|
||||
def test_plan_without_action_nudge_is_announced_on_the_status_channel(monkeypatch):
|
||||
"""The re-prompted turn is hidden, so without a badge the UI looks frozen."""
|
||||
|
||||
payloads: list[dict] = []
|
||||
backend = _make_backend(monkeypatch, _nudge_then_search_streams(), payloads)
|
||||
monkeypatch.setattr(
|
||||
"core.inference.tools.execute_tool",
|
||||
lambda *_a, **_k: "Search results: red is #f00.",
|
||||
)
|
||||
|
||||
events = list(
|
||||
backend.generate_chat_completion_with_tools(
|
||||
messages = [{"role": "user", "content": "What colour is the square?"}],
|
||||
tools = [_WEB_SEARCH_TOOL],
|
||||
max_tool_iterations = 2,
|
||||
)
|
||||
)
|
||||
|
||||
statuses = _status_texts(events)
|
||||
assert NUDGE_TOOL_CALLS_STATUS in statuses
|
||||
index = statuses.index(NUDGE_TOOL_CALLS_STATUS)
|
||||
# Blank first: the route resets its text cursor only on an empty status.
|
||||
# index > 0 matters: at 0, statuses[-1] wraps to the terminal clear.
|
||||
assert index > 0 and statuses[index - 1] == ""
|
||||
assert statuses[index + 1].startswith("Searching:")
|
||||
assert statuses[-1] == ""
|
||||
|
||||
|
||||
def test_plan_without_action_nudge_status_clears_when_the_retry_just_answers(monkeypatch):
|
||||
streams = [
|
||||
[_sse({"content": "I will search the web now."}), _done()],
|
||||
[_sse({"content": "No search needed. Final answer: the square is red."}), _done()],
|
||||
]
|
||||
payloads: list[dict] = []
|
||||
backend = _make_backend(monkeypatch, streams, payloads)
|
||||
|
||||
events = list(
|
||||
backend.generate_chat_completion_with_tools(
|
||||
messages = [{"role": "user", "content": "What colour is the square?"}],
|
||||
tools = [_WEB_SEARCH_TOOL],
|
||||
max_tool_iterations = 2,
|
||||
)
|
||||
)
|
||||
|
||||
statuses = _status_texts(events)
|
||||
assert NUDGE_TOOL_CALLS_STATUS in statuses
|
||||
assert statuses[-1] == ""
|
||||
|
||||
|
||||
def test_direct_answer_never_shows_the_nudge_status(monkeypatch):
|
||||
payloads: list[dict] = []
|
||||
backend = _make_backend(
|
||||
monkeypatch,
|
||||
[[_sse({"content": "The square is red."}), _done()]],
|
||||
payloads,
|
||||
)
|
||||
|
||||
events = list(
|
||||
backend.generate_chat_completion_with_tools(
|
||||
messages = [{"role": "user", "content": "What colour is the square?"}],
|
||||
tools = [_WEB_SEARCH_TOOL],
|
||||
max_tool_iterations = 2,
|
||||
)
|
||||
)
|
||||
|
||||
assert NUDGE_TOOL_CALLS_STATUS not in _status_texts(events)
|
||||
|
||||
|
||||
def test_nudge_status_absent_when_nudging_is_disabled(monkeypatch):
|
||||
payloads: list[dict] = []
|
||||
backend = _make_backend(monkeypatch, _nudge_then_search_streams(), payloads)
|
||||
monkeypatch.setattr(
|
||||
"core.inference.tools.execute_tool",
|
||||
lambda *_a, **_k: "Search results: red is #f00.",
|
||||
)
|
||||
|
||||
events = list(
|
||||
backend.generate_chat_completion_with_tools(
|
||||
messages = [{"role": "user", "content": "What colour is the square?"}],
|
||||
tools = [_WEB_SEARCH_TOOL],
|
||||
max_tool_iterations = 2,
|
||||
nudge_tool_calls = False,
|
||||
)
|
||||
)
|
||||
|
||||
assert NUDGE_TOOL_CALLS_STATUS not in _status_texts(events)
|
||||
assert len(payloads) == 1
|
||||
|
||||
|
||||
def test_confirm_tool_calls_allow_executes_gguf_tool(monkeypatch):
|
||||
streams = [
|
||||
_structured_tool_call("python", {"code": "print(1)"}, "call_py"),
|
||||
|
|
|
|||
|
|
@ -77,8 +77,7 @@ validate_extra_args = _lsa.validate_extra_args
|
|||
["--reasoning-format", "deepseek"],
|
||||
["-rea", "auto"],
|
||||
# Soft-managed: user flags last-wins over Unsloth's auto-set version.
|
||||
# --parallel / -np / --n-parallel are hard-denied (KV-cache + slot
|
||||
# count would desync); use `unsloth studio run --parallel N` instead.
|
||||
# --parallel / -np / --n-parallel are hard-denied; use Parallel Slots.
|
||||
["-c", "131072"],
|
||||
["--ctx-size", "8192"],
|
||||
["--flash-attn", "off"],
|
||||
|
|
@ -128,7 +127,7 @@ def test_non_flag_token_passes_through():
|
|||
@pytest.mark.parametrize(
|
||||
"denied",
|
||||
[
|
||||
# Parallel slots -- owned by the typer --parallel flag.
|
||||
# Parallel slots -- owned by typer --parallel and LoadRequest.n_parallel.
|
||||
"-np",
|
||||
"--parallel",
|
||||
"--n-parallel",
|
||||
|
|
@ -201,9 +200,8 @@ def test_denylist_rejects_all_aliases(denied):
|
|||
@pytest.mark.parametrize(
|
||||
"args,offending",
|
||||
[
|
||||
# Pass-through --parallel would last-wins-override the real slot
|
||||
# count while Unsloth's KV-cache fit + llama_parallel_slots stay at
|
||||
# the typer value -- plan vs. process disagree.
|
||||
# Pass-through --parallel would last-wins-override the real slot count
|
||||
# while the KV-cache fit and slot bookkeeping stay at the resolved value.
|
||||
(["--parallel", "8"], "--parallel"),
|
||||
(["--parallel=8"], "--parallel"),
|
||||
(["--n-parallel", "16"], "--n-parallel"),
|
||||
|
|
@ -213,7 +211,7 @@ def test_denylist_rejects_all_aliases(denied):
|
|||
# `["-np8"]` must still resolve to managed.
|
||||
(["-np8"], "-np"),
|
||||
(["-np64"], "-np"),
|
||||
# Out-of-range values that would bypass the typer 1..64 guard.
|
||||
# Out-of-range values that would bypass the PARALLEL_MIN/MAX bounds.
|
||||
(["--parallel", "999"], "--parallel"),
|
||||
(["-np", "0"], "-np"),
|
||||
(["-np999"], "-np"),
|
||||
|
|
@ -300,7 +298,7 @@ def test_is_managed_flag_true_for_denied():
|
|||
assert is_managed_flag("--api-key") is True
|
||||
assert is_managed_flag("-m") is True
|
||||
assert is_managed_flag("--model") is True
|
||||
# Parallel slots owned by the typer --parallel flag.
|
||||
# Parallel slots owned by typer --parallel and LoadRequest.n_parallel.
|
||||
assert is_managed_flag("--parallel") is True
|
||||
assert is_managed_flag("--n-parallel") is True
|
||||
assert is_managed_flag("-np") is True
|
||||
|
|
|
|||
|
|
@ -175,3 +175,46 @@ def test_call_tool_sync_passes_raise_on_error_false_and_keeps_error_images(monke
|
|||
assert out.startswith("Error: boom")
|
||||
assert MCP_IMAGES_SENTINEL in out
|
||||
assert is_tool_error(out)
|
||||
|
||||
|
||||
def test_stdio_session_call_also_passes_raise_on_error_false(monkeypatch):
|
||||
seen = {}
|
||||
|
||||
class _FakeStdioClient:
|
||||
def __init__(self):
|
||||
self.connected = False
|
||||
self.transport = SimpleNamespace(_is_session_dead = lambda: False)
|
||||
|
||||
async def __aenter__(self):
|
||||
self.connected = True
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *exc):
|
||||
self.connected = False
|
||||
|
||||
def is_connected(self):
|
||||
return self.connected
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
name,
|
||||
args,
|
||||
raise_on_error = True,
|
||||
):
|
||||
seen["raise_on_error"] = raise_on_error
|
||||
return _result(_text("boom"), _image(), is_error = True)
|
||||
|
||||
monkeypatch.setattr(
|
||||
mcp_client, "_client", lambda url, headers, use_oauth = False: _FakeStdioClient()
|
||||
)
|
||||
try:
|
||||
out = call_tool_sync(
|
||||
"npx fake-stdio-server", None, "take_screenshot", {}, scope = "s=p:t=thread1"
|
||||
)
|
||||
finally:
|
||||
mcp_client.close_stdio_sessions()
|
||||
|
||||
assert seen["raise_on_error"] is False
|
||||
assert out.startswith("Error: boom")
|
||||
assert MCP_IMAGES_SENTINEL in out
|
||||
assert is_tool_error(out)
|
||||
|
|
|
|||
|
|
@ -60,7 +60,12 @@ class FakeClient:
|
|||
def is_connected(self) -> bool:
|
||||
return self.connected
|
||||
|
||||
async def call_tool(self, name: str, args: dict):
|
||||
async def call_tool(
|
||||
self,
|
||||
name: str,
|
||||
args: dict,
|
||||
raise_on_error: bool = True,
|
||||
):
|
||||
if self.call_delay:
|
||||
await asyncio.sleep(self.call_delay)
|
||||
if self.fail_next:
|
||||
|
|
@ -120,10 +125,15 @@ def test_tool_error_does_not_recycle_session(fake_clients, monkeypatch):
|
|||
from fastmcp.exceptions import ToolError
|
||||
|
||||
class ToolFailure(FakeClient):
|
||||
async def call_tool(self, name, args):
|
||||
async def call_tool(
|
||||
self,
|
||||
name,
|
||||
args,
|
||||
raise_on_error = True,
|
||||
):
|
||||
if name == "boom":
|
||||
raise ToolError("tool exploded") # tool-level: session stays connected
|
||||
return await super().call_tool(name, args)
|
||||
return await super().call_tool(name, args, raise_on_error)
|
||||
|
||||
monkeypatch.setattr(
|
||||
mcp_client, "_client", lambda url, headers, use_oauth = False: ToolFailure(url)
|
||||
|
|
@ -441,12 +451,17 @@ def test_overlapping_calls_serialize_on_shared_session(fake_clients, monkeypatch
|
|||
active = 0
|
||||
max_active = 0
|
||||
|
||||
async def call_tool(self, name, args):
|
||||
async def call_tool(
|
||||
self,
|
||||
name,
|
||||
args,
|
||||
raise_on_error = True,
|
||||
):
|
||||
OverlapDetect.active += 1
|
||||
OverlapDetect.max_active = max(OverlapDetect.max_active, OverlapDetect.active)
|
||||
try:
|
||||
await asyncio.sleep(0.2)
|
||||
return await super().call_tool(name, args)
|
||||
return await super().call_tool(name, args, raise_on_error)
|
||||
finally:
|
||||
OverlapDetect.active -= 1
|
||||
|
||||
|
|
@ -473,9 +488,14 @@ def test_timeout_budget_spans_connect_and_call(fake_clients, monkeypatch):
|
|||
await asyncio.sleep(0.4)
|
||||
return await super().__aenter__()
|
||||
|
||||
async def call_tool(self, name, args):
|
||||
async def call_tool(
|
||||
self,
|
||||
name,
|
||||
args,
|
||||
raise_on_error = True,
|
||||
):
|
||||
await asyncio.sleep(0.5)
|
||||
return await super().call_tool(name, args)
|
||||
return await super().call_tool(name, args, raise_on_error)
|
||||
|
||||
monkeypatch.setattr(mcp_client, "_client", lambda url, headers, use_oauth = False: SlowBoth(url))
|
||||
start = time.monotonic()
|
||||
|
|
@ -565,7 +585,11 @@ def test_execute_tool_config_check_tracks_row(tmp_path, monkeypatch):
|
|||
|
||||
|
||||
def test_multi_block_result_flattens_through_session(fake_clients):
|
||||
async def _rich_call(name, args):
|
||||
async def _rich_call(
|
||||
name,
|
||||
args,
|
||||
raise_on_error = True,
|
||||
):
|
||||
return SimpleNamespace(
|
||||
content = [
|
||||
SimpleNamespace(type = "text", text = "### Page"),
|
||||
|
|
|
|||
|
|
@ -817,7 +817,26 @@ class TestExtraArgsMtpDetection:
|
|||
],
|
||||
)
|
||||
def test_flash_attn_last_value_wins(self, args, expected):
|
||||
assert _flash_attn_enabled_from_args(args) is expected
|
||||
assert _flash_attn_enabled_from_args(args, env = {}) is expected
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"value,expected",
|
||||
[
|
||||
("off", False),
|
||||
("disabled", False),
|
||||
("false", False),
|
||||
("0", False),
|
||||
("on", True),
|
||||
("auto", True),
|
||||
("garbage", True), # llama.cpp refuses to start, so the default is moot
|
||||
],
|
||||
)
|
||||
def test_flash_attn_env_applies(self, value, expected):
|
||||
env = {"LLAMA_ARG_FLASH_ATTN": value}
|
||||
assert _flash_attn_enabled_from_args([], env = env) is expected
|
||||
# llama.cpp parses the environment first, so an explicit flag still wins.
|
||||
assert _flash_attn_enabled_from_args(["-fa", "on"], env = env) is True
|
||||
assert _flash_attn_enabled_from_args(["-fa", "off"], env = env) is False
|
||||
|
||||
def test_effective_main_cache_types_follow_env_then_cli(self):
|
||||
env = {
|
||||
|
|
|
|||
517
studio/backend/tests/test_parallel_slots_per_load.py
Normal file
517
studio/backend/tests/test_parallel_slots_per_load.py
Normal file
|
|
@ -0,0 +1,517 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""Backend contract for the per-load parallel-slots knob.
|
||||
|
||||
An optional ``n_parallel`` (llama-server ``--parallel``) rides on LoadRequest;
|
||||
omitted, the server-wide launch default (``run.py --parallel``) applies. These
|
||||
tests pin the pydantic contract and the shared PARALLEL_MIN/MAX mirrors, the
|
||||
``requested_parallel_slots`` lifecycle, the ``_already_in_target_state``
|
||||
requested-vs-requested reload branch with its diffusion skip, and the route
|
||||
wiring behind the /load, /validate and /status echoes.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import re
|
||||
import struct
|
||||
import sys
|
||||
import types as _types
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
_BACKEND_DIR = str(Path(__file__).resolve().parent.parent)
|
||||
if _BACKEND_DIR not in sys.path:
|
||||
sys.path.insert(0, _BACKEND_DIR)
|
||||
|
||||
# Same external-dep stubs as the other llama_cpp unit tests.
|
||||
_loggers_stub = _types.ModuleType("loggers")
|
||||
_loggers_stub.get_logger = lambda name: __import__("logging").getLogger(name)
|
||||
sys.modules.setdefault("loggers", _loggers_stub)
|
||||
|
||||
_structlog_stub = _types.ModuleType("structlog")
|
||||
_structlog_stub.get_logger = lambda *a, **k: __import__("logging").getLogger("stub")
|
||||
sys.modules.setdefault("structlog", _structlog_stub)
|
||||
|
||||
# Real httpx: a stub would poison a combined run (routes/inference reads its
|
||||
# attrs at def time).
|
||||
import httpx # noqa: F401
|
||||
|
||||
from core.inference import llama_cpp as llama_cpp_module
|
||||
from core.inference.llama_server_args import PARALLEL_MAX, PARALLEL_MIN
|
||||
from core.inference.llama_cpp import LlamaCppBackend
|
||||
from models.inference import (
|
||||
InferenceStatusResponse,
|
||||
LoadRequest,
|
||||
LoadResponse,
|
||||
ValidateModelRequest,
|
||||
)
|
||||
|
||||
|
||||
class _FakeProcess:
|
||||
def terminate(self):
|
||||
pass
|
||||
|
||||
def wait(self, timeout = None):
|
||||
return 0
|
||||
|
||||
def kill(self):
|
||||
pass
|
||||
|
||||
def poll(self):
|
||||
return 0
|
||||
|
||||
|
||||
# ── Pydantic contract ────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_load_request_defaults_n_parallel_none():
|
||||
assert LoadRequest(model_path = "owner/repo").n_parallel is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [PARALLEL_MIN, 4, PARALLEL_MAX])
|
||||
def test_load_request_accepts_in_range_n_parallel(value):
|
||||
assert LoadRequest(model_path = "owner/repo", n_parallel = value).n_parallel == value
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [0, -1, PARALLEL_MAX + 1])
|
||||
def test_load_request_rejects_out_of_range_n_parallel(value):
|
||||
with pytest.raises(ValueError):
|
||||
LoadRequest(model_path = "owner/repo", n_parallel = value)
|
||||
|
||||
|
||||
def test_load_request_round_trips_json_key():
|
||||
req = LoadRequest.model_validate({"model_path": "owner/repo", "n_parallel": 8})
|
||||
assert req.n_parallel == 8
|
||||
assert req.model_dump()["n_parallel"] == 8
|
||||
|
||||
|
||||
def test_validate_request_n_parallel_contract():
|
||||
# /validate sizes like /load, so it carries the same field and bounds.
|
||||
assert ValidateModelRequest(model_path = "owner/repo").n_parallel is None
|
||||
assert (
|
||||
ValidateModelRequest(model_path = "owner/repo", n_parallel = PARALLEL_MAX).n_parallel
|
||||
== PARALLEL_MAX
|
||||
)
|
||||
with pytest.raises(ValueError):
|
||||
ValidateModelRequest(model_path = "owner/repo", n_parallel = PARALLEL_MAX + 1)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_cls", [LoadResponse, InferenceStatusResponse])
|
||||
def test_response_models_emit_parallel_slot_fields(model_cls):
|
||||
kwargs = (
|
||||
dict(status = "loaded", model = "owner/repo", display_name = "repo", inference = {})
|
||||
if model_cls is LoadResponse
|
||||
else {}
|
||||
)
|
||||
empty = model_cls(**kwargs).model_dump()
|
||||
assert empty["requested_parallel_slots"] is None
|
||||
assert empty["parallel_slots"] is None
|
||||
dumped = model_cls(**kwargs, requested_parallel_slots = 8, parallel_slots = 4).model_dump()
|
||||
assert dumped["requested_parallel_slots"] == 8
|
||||
assert dumped["parallel_slots"] == 4
|
||||
|
||||
|
||||
# ── Shared bounds and their deliberate mirrors ───────────────────────
|
||||
|
||||
|
||||
def _mirrored_bounds(source_path: Path) -> tuple[int, int]:
|
||||
src = source_path.read_text(encoding = "utf-8")
|
||||
low = re.search(r"^_PARALLEL_MIN\s*=\s*(\d+)$", src, re.MULTILINE)
|
||||
high = re.search(r"^_PARALLEL_MAX\s*=\s*(\d+)$", src, re.MULTILINE)
|
||||
assert low and high, f"{source_path} must define _PARALLEL_MIN/_PARALLEL_MAX"
|
||||
return int(low.group(1)), int(high.group(1))
|
||||
|
||||
|
||||
def test_run_py_mirror_matches_shared_bounds():
|
||||
assert _mirrored_bounds(Path(_BACKEND_DIR) / "run.py") == (PARALLEL_MIN, PARALLEL_MAX)
|
||||
|
||||
|
||||
def test_cli_mirror_matches_shared_bounds():
|
||||
cli = Path(_BACKEND_DIR).parent.parent / "unsloth_cli" / "commands" / "studio.py"
|
||||
assert _mirrored_bounds(cli) == (PARALLEL_MIN, PARALLEL_MAX)
|
||||
|
||||
|
||||
def test_frontend_mirror_matches_shared_bounds():
|
||||
# The UI clamps with its own copy; a bumped PARALLEL_MAX that skips it would
|
||||
# leave the UI silently capping lower.
|
||||
src = (
|
||||
Path(_BACKEND_DIR).parent
|
||||
/ "frontend"
|
||||
/ "src"
|
||||
/ "features"
|
||||
/ "model-picker"
|
||||
/ "model-config"
|
||||
/ "per-model-config.ts"
|
||||
).read_text(encoding = "utf-8")
|
||||
low = re.search(r"^export const N_PARALLEL_MIN = (\d+);$", src, re.MULTILINE)
|
||||
high = re.search(r"^export const N_PARALLEL_MAX = (\d+);$", src, re.MULTILINE)
|
||||
assert low and high, "per-model-config.ts must export N_PARALLEL_MIN/MAX"
|
||||
assert (int(low.group(1)), int(high.group(1))) == (PARALLEL_MIN, PARALLEL_MAX)
|
||||
|
||||
|
||||
def test_preset_model_reuses_shared_bounds():
|
||||
# Bounds drifting from PARALLEL_MIN/MAX would 422 valid presets on every sync.
|
||||
from routes.chat_history import ChatPresetLoadConfig
|
||||
|
||||
field = ChatPresetLoadConfig.model_fields["nParallel"]
|
||||
bounds = {type(m).__name__: getattr(m, "ge", getattr(m, "le", None)) for m in field.metadata}
|
||||
assert bounds.get("Ge") == PARALLEL_MIN
|
||||
assert bounds.get("Le") == PARALLEL_MAX
|
||||
|
||||
|
||||
# ── requested_parallel_slots lifecycle ───────────────────────────────
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def backend(monkeypatch):
|
||||
monkeypatch.setattr(LlamaCppBackend, "_kill_orphaned_servers", lambda self: 0)
|
||||
monkeypatch.setattr(llama_cpp_module.atexit, "register", lambda *_args, **_kwargs: None)
|
||||
return LlamaCppBackend()
|
||||
|
||||
|
||||
def test_requested_parallel_slots_initial_value_is_one(backend):
|
||||
assert backend.requested_parallel_slots == 1
|
||||
|
||||
|
||||
def test_requested_parallel_slots_reflects_field(backend):
|
||||
backend._requested_n_parallel = 8
|
||||
assert backend.requested_parallel_slots == 8
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [None, 0, -2, "not-an-int"])
|
||||
def test_requested_parallel_slots_invalid_value_falls_back_to_one(backend, value):
|
||||
backend._requested_n_parallel = value
|
||||
assert backend.requested_parallel_slots == 1
|
||||
|
||||
|
||||
def test_reset_effective_parallel_slots_also_resets_requested(backend):
|
||||
backend._requested_n_parallel = 8
|
||||
backend._commit_effective_parallel_slots(4)
|
||||
|
||||
backend._reset_effective_parallel_slots()
|
||||
|
||||
assert backend.requested_parallel_slots == 1
|
||||
assert backend.effective_parallel_slots == 1
|
||||
|
||||
|
||||
def test_unload_resets_requested_parallel_slots(backend):
|
||||
backend._process = _FakeProcess()
|
||||
backend._requested_n_parallel = 8
|
||||
|
||||
backend.unload_model()
|
||||
|
||||
assert backend.requested_parallel_slots == 1
|
||||
|
||||
|
||||
def test_load_model_commits_requested_from_pending_kwargs():
|
||||
# n_parallel may be reduced before the commit, so the requested value must
|
||||
# come from the pre-reduction pending snapshot.
|
||||
src = inspect.getsource(LlamaCppBackend.load_model)
|
||||
commit = src.find(
|
||||
'self._requested_n_parallel = max(1, int(_pending_load_kwargs["n_parallel"]))'
|
||||
)
|
||||
healthy = src.find("self._healthy = True\n", 0, commit if commit != -1 else None)
|
||||
snapshot = src.find("self._last_load_kwargs = _pending_load_kwargs")
|
||||
assert commit != -1, "load_model must commit the requested slot count"
|
||||
assert healthy != -1 and healthy < commit < snapshot
|
||||
|
||||
|
||||
# ── _already_in_target_state requested-vs-requested branch ───────────
|
||||
|
||||
|
||||
def _loaded_backend() -> LlamaCppBackend:
|
||||
backend = LlamaCppBackend()
|
||||
backend._process = _FakeProcess() # is_loaded only checks "is not None"
|
||||
backend._healthy = True
|
||||
backend._model_identifier = "owner/repo"
|
||||
backend._hf_variant = "Q4_K_M"
|
||||
backend._requested_n_ctx = 8192
|
||||
backend._cache_type_kv = None
|
||||
backend._requested_spec_mode = "auto"
|
||||
backend._chat_template_override = None
|
||||
backend._is_vision = False
|
||||
backend._extra_args = None
|
||||
backend._gguf_path = None
|
||||
return backend
|
||||
|
||||
|
||||
def _target_state(backend: LlamaCppBackend, n_parallel: int) -> bool:
|
||||
return backend._already_in_target_state(
|
||||
gguf_path = None,
|
||||
model_identifier = "owner/repo",
|
||||
hf_variant = "Q4_K_M",
|
||||
n_ctx = 8192,
|
||||
cache_type_kv = None,
|
||||
speculative_type = "auto",
|
||||
chat_template_override = None,
|
||||
extra_args = None,
|
||||
is_vision = False,
|
||||
n_parallel = n_parallel,
|
||||
)
|
||||
|
||||
|
||||
def test_already_in_target_state_matches_same_slots():
|
||||
backend = _loaded_backend()
|
||||
backend._requested_n_parallel = 4
|
||||
assert _target_state(backend, 4) is True
|
||||
|
||||
|
||||
def test_already_in_target_state_reloads_on_slots_change():
|
||||
backend = _loaded_backend()
|
||||
backend._requested_n_parallel = 4
|
||||
assert _target_state(backend, 8) is False
|
||||
|
||||
|
||||
def test_already_in_target_state_compares_requested_not_effective():
|
||||
# An identical re-Apply must dedupe even after the fitter reduced the slots.
|
||||
backend = _loaded_backend()
|
||||
backend._requested_n_parallel = 8
|
||||
backend._commit_effective_parallel_slots(4)
|
||||
assert _target_state(backend, 8) is True
|
||||
|
||||
|
||||
def test_already_in_target_state_ignores_slots_for_diffusion():
|
||||
# The diffusion runner ignores --parallel, so a slots change must not reload.
|
||||
backend = _loaded_backend()
|
||||
backend._is_diffusion = True
|
||||
backend._requested_n_parallel = 1
|
||||
assert _target_state(backend, 8) is True
|
||||
|
||||
|
||||
# ── Route wiring (source contract, mirroring test_gpu_memory_mode) ───
|
||||
|
||||
|
||||
def _route_source() -> str:
|
||||
return (Path(_BACKEND_DIR) / "routes" / "inference.py").read_text(encoding = "utf-8")
|
||||
|
||||
|
||||
def _load_impl_source() -> str:
|
||||
"""Body of _load_model_impl only, so positional assertions can't be
|
||||
satisfied by a later function in the module."""
|
||||
src = _route_source()
|
||||
body = src[src.index("async def _load_model_impl") :]
|
||||
return body[: body.index("\n@router.")]
|
||||
|
||||
|
||||
def test_route_resolves_slots_once_before_dedupe_guard_and_load():
|
||||
load_impl = _load_impl_source()
|
||||
resolve = load_impl.index("request.n_parallel")
|
||||
fallback = load_impl.index('getattr(_app_state, "llama_parallel_slots", 1)')
|
||||
dedupe = load_impl.index("requested_parallel_slots = _n_parallel")
|
||||
guard = load_impl.index("_guard_chat_load_against_training")
|
||||
# The GGUF launch kwargs, not the guard's own kwarg (which shares the spelling).
|
||||
load_kwargs = load_impl.index("_common_load_kwargs = dict(")
|
||||
assert resolve < dedupe, "resolution must precede the reload dedupe"
|
||||
assert fallback < dedupe
|
||||
assert resolve < guard < load_kwargs
|
||||
# Guard and load kwargs share the resolved value; app.state is read once.
|
||||
assert load_impl.count("n_parallel = _n_parallel") == 2
|
||||
assert "n_parallel = _n_parallel" in load_impl[load_kwargs : load_kwargs + 800]
|
||||
assert load_impl.count('getattr(_app_state, "llama_parallel_slots", 1)') == 1
|
||||
# getattr, so a direct caller without an app cannot raise, and no re-read.
|
||||
assert "fastapi_request.app.state" not in load_impl
|
||||
|
||||
|
||||
def test_route_dedupe_compares_requested_slots_and_skips_diffusion():
|
||||
match_impl = _route_source()[_route_source().index("def _request_matches_loaded_settings") :]
|
||||
match_impl = match_impl[: match_impl.index("\ndef ")]
|
||||
assert "requested_parallel_slots is not None" in match_impl
|
||||
assert "not llama_backend.is_diffusion" in match_impl
|
||||
assert "llama_backend.requested_parallel_slots" in match_impl
|
||||
|
||||
|
||||
def test_route_echoes_requested_and_effective_slots():
|
||||
route_src = _route_source()
|
||||
# Both /load returns plus the /status GGUF branch, via the shared helper.
|
||||
assert route_src.count("**_parallel_slot_echo(llama_backend)") == 3
|
||||
|
||||
|
||||
def test_parallel_slot_echo_reports_none_for_diffusion():
|
||||
# Diffusion never commits a count, so echoing the reset placeholder 1 would lie.
|
||||
from routes.inference import _parallel_slot_echo
|
||||
|
||||
backend = _loaded_backend()
|
||||
backend._requested_n_parallel = 8
|
||||
backend._commit_effective_parallel_slots(4)
|
||||
assert _parallel_slot_echo(backend) == {"requested_parallel_slots": 8, "parallel_slots": 4}
|
||||
backend._is_diffusion = True
|
||||
assert _parallel_slot_echo(backend) == {
|
||||
"requested_parallel_slots": None,
|
||||
"parallel_slots": None,
|
||||
}
|
||||
|
||||
|
||||
def test_validate_route_prefers_request_n_parallel():
|
||||
validate_impl = _route_source()[_route_source().index("async def validate_model") :]
|
||||
resolve = validate_impl.index("request.n_parallel")
|
||||
fallback = validate_impl.index('"llama_parallel_slots",')
|
||||
guard = validate_impl.index("_guard_chat_load_against_training")
|
||||
assert guard < resolve and guard < fallback, "the guard call resolves the slots inline"
|
||||
|
||||
|
||||
def _load_model_source() -> str:
|
||||
return inspect.getsource(LlamaCppBackend.load_model)
|
||||
|
||||
|
||||
def test_slots_fall_back_to_one_without_kv_unified():
|
||||
# Without --kv-unified llama-server gives each slot -c/N, so an explicit
|
||||
# --parallel N shrinks every context window.
|
||||
src = _load_model_source()
|
||||
clamp = src.find("supports_kv_unified")
|
||||
assert clamp != -1, "load_model must check for --kv-unified before honouring the slots"
|
||||
block = src[clamp : clamp + 700]
|
||||
assert (
|
||||
"n_parallel > 1" in src[clamp - 300 : clamp]
|
||||
), "only an explicit multi-slot load is clamped"
|
||||
assert "n_parallel = 1" in block
|
||||
|
||||
|
||||
def test_clamp_sits_between_the_echo_and_the_fit():
|
||||
# The echo reports the ask and the fit uses what launches, so the clamp
|
||||
# belongs between the two.
|
||||
src = _load_model_source()
|
||||
pending = src.index("_pending_load_kwargs")
|
||||
clamp = src.index("supports_kv_unified")
|
||||
estimate = src.index("_estimate")
|
||||
commit = src.index("_commit_effective_parallel_slots")
|
||||
assert pending < clamp, "the requested count is captured before the clamp"
|
||||
assert clamp < estimate, "the fit must be estimated from the effective slot count"
|
||||
assert clamp < commit, "the committed effective count is the clamped one"
|
||||
|
||||
|
||||
# ── Training-guard sizing ────────────────────────────────────────────
|
||||
|
||||
|
||||
def _write_swa_gguf(path: Path) -> str:
|
||||
"""Smallest DiffusionGemma-shaped header the KV estimator can size: the
|
||||
canvas marker routing it to the diffusion runner, plus the sliding-window
|
||||
dims that make llama.cpp's SWA cache slot-scaled."""
|
||||
|
||||
def _kv_str(key: str, value: str) -> bytes:
|
||||
kb, vb = key.encode(), value.encode()
|
||||
return (
|
||||
struct.pack("<Q", len(kb)) + kb + struct.pack("<I", 8) + struct.pack("<Q", len(vb)) + vb
|
||||
)
|
||||
|
||||
def _kv_u32(key: str, value: int) -> bytes:
|
||||
kb = key.encode()
|
||||
return struct.pack("<Q", len(kb)) + kb + struct.pack("<I", 4) + struct.pack("<I", value)
|
||||
|
||||
arch = "diffusion-gemma"
|
||||
kvs = [
|
||||
_kv_str("general.architecture", arch),
|
||||
_kv_u32("diffusion.canvas_length", 256),
|
||||
_kv_u32(f"{arch}.context_length", 32768),
|
||||
_kv_u32(f"{arch}.block_count", 30),
|
||||
_kv_u32(f"{arch}.attention.head_count", 16),
|
||||
_kv_u32(f"{arch}.attention.head_count_kv", 8),
|
||||
_kv_u32(f"{arch}.attention.key_length", 512),
|
||||
_kv_u32(f"{arch}.attention.value_length", 512),
|
||||
_kv_u32(f"{arch}.attention.sliding_window", 1024),
|
||||
_kv_u32(f"{arch}.attention.key_length_swa", 256),
|
||||
_kv_u32(f"{arch}.attention.value_length_swa", 256),
|
||||
]
|
||||
path.write_bytes(struct.pack("<IIQQ", 0x46554747, 3, 0, len(kvs)) + b"".join(kvs))
|
||||
return str(path)
|
||||
|
||||
|
||||
def _guard_required_gb(
|
||||
monkeypatch,
|
||||
gguf_path: str,
|
||||
*,
|
||||
n_parallel: int,
|
||||
diffusion,
|
||||
caps = None,
|
||||
) -> float:
|
||||
"""Run the training guard over a local GGUF and return the size it budgeted."""
|
||||
import routes.inference as inf
|
||||
|
||||
seen = {}
|
||||
|
||||
core_training = _types.ModuleType("core.training")
|
||||
core_training.get_training_backend = lambda: _types.SimpleNamespace(
|
||||
is_training_active = lambda: True
|
||||
)
|
||||
|
||||
def _can_load(**kwargs):
|
||||
seen.update(kwargs)
|
||||
return True, {"mode": "single_device"}
|
||||
|
||||
training_vram = _types.ModuleType("routes.training_vram")
|
||||
training_vram.can_load_chat_during_training = _can_load
|
||||
monkeypatch.setitem(sys.modules, "core.training", core_training)
|
||||
monkeypatch.setitem(sys.modules, "routes.training_vram", training_vram)
|
||||
|
||||
monkeypatch.setattr(inf, "_classify_diffusion_gguf", lambda _config: diffusion)
|
||||
monkeypatch.setattr(LlamaCppBackend, "_is_vulkan_backend", staticmethod(lambda *a, **k: False))
|
||||
monkeypatch.setattr(LlamaCppBackend, "_effective_gpu_count", staticmethod(lambda *a, **k: 1))
|
||||
monkeypatch.setattr(LlamaCppBackend, "_diffusion_gpu_arg", staticmethod(lambda *a, **k: "0"))
|
||||
# Pin the --kv-unified probe so the estimate cannot depend on a locally
|
||||
# installed llama-server. Default "no binary found" leaves the count alone.
|
||||
monkeypatch.setattr(
|
||||
LlamaCppBackend,
|
||||
"probe_server_capabilities",
|
||||
classmethod(lambda cls, binary = None: dict(caps or {})),
|
||||
)
|
||||
|
||||
inf._guard_chat_load_against_training(
|
||||
_types.SimpleNamespace(is_gguf = True, gguf_file = gguf_path, identifier = "local/model"),
|
||||
model_identifier = "local/model",
|
||||
hf_token = None,
|
||||
load_in_4bit = False,
|
||||
max_seq_length = 8192,
|
||||
requested_gpu_ids = None,
|
||||
n_parallel = n_parallel,
|
||||
gpu_memory_mode = "auto",
|
||||
)
|
||||
return seen["required_override_gb"]
|
||||
|
||||
|
||||
def test_training_guard_sizes_a_diffusion_gguf_at_one_slot(monkeypatch, tmp_path):
|
||||
# Diffusion ignores --parallel, so slots must not inflate the estimate and 409
|
||||
# a load that would have fitted beside training.
|
||||
gguf = _write_swa_gguf(tmp_path / "diffusion.gguf")
|
||||
one = _guard_required_gb(monkeypatch, gguf, n_parallel = 1, diffusion = True)
|
||||
many = _guard_required_gb(monkeypatch, gguf, n_parallel = 8, diffusion = True)
|
||||
assert one == many
|
||||
|
||||
|
||||
def test_training_guard_still_sizes_slots_for_an_ordinary_gguf(monkeypatch, tmp_path):
|
||||
# llama-server does allocate per-slot SWA cells, so the reduction above must
|
||||
# be scoped to diffusion and not flatten every GGUF to one slot.
|
||||
gguf = _write_swa_gguf(tmp_path / "chat.gguf")
|
||||
one = _guard_required_gb(monkeypatch, gguf, n_parallel = 1, diffusion = False)
|
||||
many = _guard_required_gb(monkeypatch, gguf, n_parallel = 8, diffusion = False)
|
||||
assert many > one
|
||||
|
||||
|
||||
def test_training_guard_sizes_one_slot_when_the_binary_has_no_kv_unified(monkeypatch, tmp_path):
|
||||
# load_model clamps a multi-slot request to 1 on such a build, where each slot
|
||||
# carries its own SWA stream, so sizing the asked count would 409 a load that fits.
|
||||
gguf = _write_swa_gguf(tmp_path / "chat.gguf")
|
||||
old = {"found": True, "supports_kv_unified": False}
|
||||
one = _guard_required_gb(monkeypatch, gguf, n_parallel = 1, diffusion = False, caps = old)
|
||||
many = _guard_required_gb(monkeypatch, gguf, n_parallel = 8, diffusion = False, caps = old)
|
||||
assert one == many
|
||||
|
||||
|
||||
def test_training_guard_sizes_every_slot_when_kv_unified_exists(monkeypatch, tmp_path):
|
||||
# The clamp is scoped to binaries that cannot serve the slots; a capable one
|
||||
# really does allocate the SWA window per slot.
|
||||
gguf = _write_swa_gguf(tmp_path / "chat.gguf")
|
||||
new = {"found": True, "supports_kv_unified": True}
|
||||
one = _guard_required_gb(monkeypatch, gguf, n_parallel = 1, diffusion = False, caps = new)
|
||||
many = _guard_required_gb(monkeypatch, gguf, n_parallel = 8, diffusion = False, caps = new)
|
||||
assert many > one
|
||||
|
||||
|
||||
def test_training_guard_keeps_slots_for_an_unclassified_gguf(monkeypatch, tmp_path):
|
||||
# None = inconclusive header, so keep the larger estimate rather than
|
||||
# under-size against training.
|
||||
gguf = _write_swa_gguf(tmp_path / "unknown.gguf")
|
||||
one = _guard_required_gb(monkeypatch, gguf, n_parallel = 1, diffusion = None)
|
||||
many = _guard_required_gb(monkeypatch, gguf, n_parallel = 8, diffusion = None)
|
||||
assert many > one
|
||||
|
|
@ -24,6 +24,7 @@ from core.inference.safetensors_agentic import (
|
|||
strip_tool_markup_streaming,
|
||||
)
|
||||
from core.inference.tool_call_parser import (
|
||||
NUDGE_TOOL_CALLS_STATUS,
|
||||
RAG_MAX_SEARCHES_PER_TURN,
|
||||
has_tool_signal,
|
||||
parse_tool_calls_from_text,
|
||||
|
|
@ -2231,6 +2232,24 @@ def test_reprompt_names_only_active_tools_not_hardcoded():
|
|||
assert "python" not in reprompt["content"]
|
||||
|
||||
|
||||
def test_reprompt_is_announced_on_the_status_channel():
|
||||
# The re-prompted turn is hidden, so the badge is the only sign of life.
|
||||
# Blank still comes first: the route resets its text cursor only on that.
|
||||
_captured, events = _reprompt_loop(auto_heal_tool_calls = True)
|
||||
statuses = [e["text"] for e in events if e["type"] == "status"]
|
||||
assert NUDGE_TOOL_CALLS_STATUS in statuses
|
||||
index = statuses.index(NUDGE_TOOL_CALLS_STATUS)
|
||||
# index > 0 matters: at 0, statuses[-1] wraps to the terminal clear.
|
||||
assert index > 0 and statuses[index - 1] == ""
|
||||
assert statuses[-1] == ""
|
||||
|
||||
|
||||
def test_reprompt_status_absent_without_a_nudge():
|
||||
_captured, events = _reprompt_loop(auto_heal_tool_calls = False)
|
||||
statuses = [e["text"] for e in events if e["type"] == "status"]
|
||||
assert NUDGE_TOOL_CALLS_STATUS not in statuses
|
||||
|
||||
|
||||
def test_reprompt_suppressed_when_auto_heal_disabled():
|
||||
# With Auto-Heal off the safetensors nudge must stay silent for backend parity
|
||||
# with the GGUF loop, so only the single initial generation runs.
|
||||
|
|
|
|||
10
studio/frontend/package-lock.json
generated
10
studio/frontend/package-lock.json
generated
|
|
@ -34,6 +34,7 @@
|
|||
"@tanstack/react-virtual": "3.13.25",
|
||||
"@tauri-apps/api": "^2.10.1",
|
||||
"@tauri-apps/plugin-clipboard-manager": "^2.3.2",
|
||||
"@tauri-apps/plugin-deep-link": "2.4.9",
|
||||
"@tauri-apps/plugin-notification": "^2.3.3",
|
||||
"@tauri-apps/plugin-opener": "^2.5.3",
|
||||
"@tauri-apps/plugin-process": "^2.3.1",
|
||||
|
|
@ -6451,6 +6452,15 @@
|
|||
"@tauri-apps/api": "^2.8.0"
|
||||
}
|
||||
},
|
||||
"node_modules/@tauri-apps/plugin-deep-link": {
|
||||
"version": "2.4.9",
|
||||
"resolved": "https://registry.npmjs.org/@tauri-apps/plugin-deep-link/-/plugin-deep-link-2.4.9.tgz",
|
||||
"integrity": "sha512-u0SKOUHnJ1wqeqXsDFq2+kASCBj9xxbG0g9XZWPy9SOmU4wXtp6b/wiYpm6oH6/5fBTQsLqnLhIvqLBRpgHJlA==",
|
||||
"license": "MIT OR Apache-2.0",
|
||||
"dependencies": {
|
||||
"@tauri-apps/api": "^2.11.0"
|
||||
}
|
||||
},
|
||||
"node_modules/@tauri-apps/plugin-notification": {
|
||||
"version": "2.3.3",
|
||||
"resolved": "https://registry.npmjs.org/@tauri-apps/plugin-notification/-/plugin-notification-2.3.3.tgz",
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@
|
|||
"@tanstack/react-virtual": "3.13.25",
|
||||
"@tauri-apps/api": "^2.10.1",
|
||||
"@tauri-apps/plugin-clipboard-manager": "^2.3.2",
|
||||
"@tauri-apps/plugin-deep-link": "2.4.9",
|
||||
"@tauri-apps/plugin-notification": "^2.3.3",
|
||||
"@tauri-apps/plugin-opener": "^2.5.3",
|
||||
"@tauri-apps/plugin-process": "^2.3.1",
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import { TooltipProvider } from "@/components/ui/tooltip";
|
|||
import { WebUpdateBanner } from "@/components/web/update-banner";
|
||||
import { fetchDeviceType } from "@/config/env";
|
||||
import { getTauriAuthFailure, tauriAutoAuth } from "@/features/auth";
|
||||
import { DeepLinkHandler } from "@/features/deep-links";
|
||||
import { DownloadManagerPanel } from "@/features/hub/download-manager";
|
||||
import { NativeIntentDrain } from "@/features/native-intents/native-intent-drain";
|
||||
import {
|
||||
|
|
@ -255,7 +256,7 @@ const MAC_NATIVE_CHROME_STYLE = {
|
|||
"--studio-non-chat-content-top-inset": "34px",
|
||||
"--studio-hidden-route-top-inset": "34px",
|
||||
"--studio-chat-header-height": "44px",
|
||||
"--studio-chat-header-padding-top": "8px",
|
||||
"--studio-chat-header-padding-top": "7px",
|
||||
"--studio-chat-control-height": "33px",
|
||||
"--studio-chat-header-right-inset": "0px",
|
||||
} as CSSProperties;
|
||||
|
|
@ -500,6 +501,7 @@ export function AppProvider({ children }: AppProviderProps) {
|
|||
<MotionConfig reducedMotion={REDUCED_MOTION_MAP[reduceMotion]}>
|
||||
<TooltipProvider>
|
||||
<AppearanceCustomizationEffect />
|
||||
<DeepLinkHandler />
|
||||
<TauriWrapper>{children}</TauriWrapper>
|
||||
<Toaster
|
||||
position="top-right"
|
||||
|
|
|
|||
|
|
@ -13,6 +13,9 @@ const ModelsPage = lazyRouteComponent(
|
|||
export interface ModelsSearch {
|
||||
tab?: "discover" | "downloaded";
|
||||
model?: string;
|
||||
file?: string;
|
||||
|
||||
intent?: number;
|
||||
section?: "trending" | "latest" | "finetune";
|
||||
kind?: "models" | "datasets";
|
||||
}
|
||||
|
|
@ -28,6 +31,18 @@ export const Route = createRoute({
|
|||
if (raw === "discover" || raw === "downloaded") next.tab = raw;
|
||||
const model = search.model;
|
||||
if (typeof model === "string" && model.length > 0) next.model = model;
|
||||
const file = search.file;
|
||||
if (next.model && typeof file === "string" && file.length > 0)
|
||||
next.file = file;
|
||||
|
||||
const intent = search.intent;
|
||||
if (
|
||||
next.file &&
|
||||
typeof intent === "number" &&
|
||||
Number.isSafeInteger(intent)
|
||||
) {
|
||||
next.intent = intent;
|
||||
}
|
||||
const section = search.section;
|
||||
if (
|
||||
section === "trending" ||
|
||||
|
|
|
|||
|
|
@ -1183,7 +1183,11 @@ export function AppSidebar() {
|
|||
<Sidebar
|
||||
collapsible="icon"
|
||||
variant="sidebar"
|
||||
className="font-heading group-data-[collapsible=icon]:[&_[data-sidebar=sidebar]]:bg-white dark:group-data-[collapsible=icon]:[&_[data-sidebar=sidebar]]:bg-background"
|
||||
className={cn(
|
||||
"font-heading group-data-[collapsible=icon]:[&_[data-sidebar=sidebar]]:bg-white dark:group-data-[collapsible=icon]:[&_[data-sidebar=sidebar]]:bg-background",
|
||||
usesNativeMacTitlebar &&
|
||||
"group-data-[collapsible=icon]:[&_[data-sidebar=sidebar]]:border-r-0",
|
||||
)}
|
||||
>
|
||||
<SidebarHeader
|
||||
className={cn(
|
||||
|
|
@ -1205,6 +1209,7 @@ export function AppSidebar() {
|
|||
/>
|
||||
)}
|
||||
<div
|
||||
data-tauri-drag-region={usesNativeMacTitlebar || undefined}
|
||||
className={cn(
|
||||
"relative z-10 flex items-center gap-[8.5px] group-data-[collapsible=icon]:hidden",
|
||||
showCompactMacBrand &&
|
||||
|
|
|
|||
|
|
@ -47,20 +47,151 @@ const normalizeLanguage = (language: string): BundledLanguage => {
|
|||
return (override ?? (key as BundledLanguage));
|
||||
};
|
||||
|
||||
// A streaming fence re-enters highlight() every frame with the whole block, so
|
||||
// Shiki re-tokenizes it in full ~60x/sec. Past MIN_INCREMENTAL_CHARS, reuse the
|
||||
// cached tokens with an unstyled tail, re-tokenizing at most every REFRESH_MS.
|
||||
const MIN_INCREMENTAL_CHARS = 2000;
|
||||
const REFRESH_MS = 250;
|
||||
// Wall-clock Date.now() can step backwards (NTP, sleep resume) and make
|
||||
// `elapsed` negative; the throttle only needs elapsed time, so stay monotonic.
|
||||
const monotonicNow = (): number =>
|
||||
typeof performance !== "undefined" && typeof performance.now === "function"
|
||||
? performance.now()
|
||||
: Date.now();
|
||||
|
||||
// One slot per fence: a message can hold several large fences, and Streamdown
|
||||
// revisits all of them on every render.
|
||||
const MAX_SLOTS_PER_KEY = 8;
|
||||
|
||||
type TokenLine = HighlightResult["tokens"][number];
|
||||
type Dispatch = {
|
||||
opts: HighlightOptions;
|
||||
language: BundledLanguage;
|
||||
callback?: (result: HighlightResult) => void;
|
||||
};
|
||||
type Slot = {
|
||||
/** Code that produced `result`. Only ever set together with it. */
|
||||
code: string;
|
||||
result: HighlightResult | null;
|
||||
/** Code of the dispatch awaiting a callback. */
|
||||
inFlight: string | null;
|
||||
lastDispatchAt: number;
|
||||
trailing: ReturnType<typeof setTimeout> | null;
|
||||
pending: Dispatch | null;
|
||||
};
|
||||
|
||||
// No colour fields, so it renders in the default foreground instead of
|
||||
// inheriting a neighbouring token's colour.
|
||||
const plainLine = (text: string): TokenLine =>
|
||||
[{ content: text, offset: 0 }] as unknown as TokenLine;
|
||||
|
||||
export function createCodePlugin(
|
||||
options: CodePluginOptions = {},
|
||||
): CodeHighlighterPlugin {
|
||||
const inner = createShikiCodePlugin(options);
|
||||
const slotsByKey = new Map<string, Slot[]>();
|
||||
|
||||
const clearTrailing = (slot: Slot) => {
|
||||
if (slot.trailing !== null) clearTimeout(slot.trailing);
|
||||
slot.trailing = null;
|
||||
slot.pending = null;
|
||||
};
|
||||
|
||||
const adopt = (slot: Slot, code: string, result: HighlightResult) => {
|
||||
// Write code and result together so a reuse cannot slice one against the other.
|
||||
slot.code = code;
|
||||
slot.result = result;
|
||||
slot.inFlight = null;
|
||||
};
|
||||
|
||||
const dispatch = (slot: Slot, d: Dispatch) => {
|
||||
slot.inFlight = d.opts.code;
|
||||
slot.lastDispatchAt = monotonicNow();
|
||||
const immediate = inner.highlight({ ...d.opts, language: d.language }, (result) => {
|
||||
if (slot.inFlight === d.opts.code) {
|
||||
adopt(slot, d.opts.code, result);
|
||||
}
|
||||
d.callback?.(result);
|
||||
});
|
||||
// @streamdown/code answers out of its own cache synchronously and never
|
||||
// invokes the callback, so adopt here too or the slot keeps older tokens.
|
||||
if (immediate) {
|
||||
adopt(slot, d.opts.code, immediate);
|
||||
}
|
||||
return immediate;
|
||||
};
|
||||
|
||||
return {
|
||||
...inner,
|
||||
supportsLanguage: (language) => inner.supportsLanguage(normalizeLanguage(language)),
|
||||
supportsLanguage: (language) =>
|
||||
inner.supportsLanguage(normalizeLanguage(language)),
|
||||
highlight: (
|
||||
opts: HighlightOptions,
|
||||
callback?: (result: HighlightResult) => void,
|
||||
) =>
|
||||
inner.highlight(
|
||||
{ ...opts, language: normalizeLanguage(opts.language) },
|
||||
callback,
|
||||
),
|
||||
) => {
|
||||
const language = normalizeLanguage(opts.language);
|
||||
if (opts.code.length < MIN_INCREMENTAL_CHARS) {
|
||||
return inner.highlight({ ...opts, language }, callback);
|
||||
}
|
||||
|
||||
const key = `${language} ${JSON.stringify(opts.themes)}`;
|
||||
let slots = slotsByKey.get(key);
|
||||
if (!slots) {
|
||||
slots = [];
|
||||
slotsByKey.set(key, slots);
|
||||
}
|
||||
|
||||
// Longest-prefix match, so sibling fences do not evict each other.
|
||||
let slot: Slot | null = null;
|
||||
let bestLength = -1;
|
||||
for (const candidate of slots) {
|
||||
const anchor = candidate.code || candidate.inFlight || "";
|
||||
if (!anchor || !opts.code.startsWith(anchor)) continue;
|
||||
if (anchor.length > bestLength) {
|
||||
slot = candidate;
|
||||
bestLength = anchor.length;
|
||||
}
|
||||
}
|
||||
if (!slot) {
|
||||
slot = { code: "", result: null, inFlight: null, lastDispatchAt: 0, trailing: null, pending: null };
|
||||
slots.unshift(slot);
|
||||
for (const dropped of slots.splice(MAX_SLOTS_PER_KEY)) clearTrailing(dropped);
|
||||
}
|
||||
|
||||
// Finished fence re-rendered unchanged: serve it, never re-tokenize.
|
||||
if (slot.result && slot.code === opts.code) return slot.result;
|
||||
|
||||
const elapsed = monotonicNow() - slot.lastDispatchAt;
|
||||
const grew = slot.result !== null && opts.code.length > slot.code.length;
|
||||
if (!grew || elapsed >= REFRESH_MS) {
|
||||
clearTrailing(slot);
|
||||
return dispatch(slot, { opts, language, callback });
|
||||
}
|
||||
|
||||
// Close out a reused run, so a final render is never left unstyled.
|
||||
slot.pending = { opts, language, callback };
|
||||
if (slot.trailing === null) {
|
||||
const target = slot;
|
||||
target.trailing = setTimeout(() => {
|
||||
target.trailing = null;
|
||||
const next = target.pending;
|
||||
target.pending = null;
|
||||
if (!next) return;
|
||||
const immediate = dispatch(target, next);
|
||||
// Nothing consumes this return value, so hand a synchronous cache
|
||||
// hit to the callback or the fence keeps its unstyled tail.
|
||||
if (immediate) next.callback?.(immediate);
|
||||
}, Math.max(0, REFRESH_MS - elapsed));
|
||||
}
|
||||
|
||||
const previous = slot.result as HighlightResult;
|
||||
// Drop the cached final line: it may have been cut mid-token.
|
||||
const keptLines = previous.tokens.slice(
|
||||
0,
|
||||
Math.max(0, slot.code.split("\n").length - 1),
|
||||
);
|
||||
const tail = opts.code.split("\n").slice(keptLines.length);
|
||||
return { ...previous, tokens: [...keptLines, ...tail.map(plainLine)] };
|
||||
},
|
||||
};
|
||||
}
|
||||
|
|
|
|||
|
|
@ -91,6 +91,7 @@ import {
|
|||
useResearchRunStore,
|
||||
} from "@/features/chat/stores/research-run-store";
|
||||
import { parseExternalModelId } from "@/features/chat/external-providers";
|
||||
import { toolStatusKind } from "@/features/chat/utils/tool-status";
|
||||
import { McpComposerButton } from "@/features/chat/mcp-composer-button";
|
||||
import { getExternalReasoningCapabilities } from "@/features/chat/provider-capabilities";
|
||||
import { useRagToolDisabled } from "@/features/chat/hooks/use-rag-tool-disabled";
|
||||
|
|
@ -2847,15 +2848,28 @@ const ToolStatusDisplay: FC = () => {
|
|||
}
|
||||
// From the store's start time, so returning to the conversation resumes rather than restarting.
|
||||
const elapsed = Math.max(0, Math.floor((now - startedAt) / 1000));
|
||||
const isRunning = toolStatus.startsWith("Running");
|
||||
const StatusIcon = isRunning ? TerminalIcon : GlobeIcon;
|
||||
const kind = toolStatusKind(toolStatus);
|
||||
const isNudging = kind === "nudge";
|
||||
const StatusIcon = kind === "terminal" ? TerminalIcon : GlobeIcon;
|
||||
return (
|
||||
<div
|
||||
data-testid="composer-tool-status"
|
||||
className="mb-2 flex w-full flex-row items-center gap-2 px-1.5 pt-0.5 pb-1"
|
||||
>
|
||||
<div className="flex animate-pulse items-center gap-2 rounded-full border border-primary/20 bg-primary/5 px-3 py-1.5 text-xs text-primary">
|
||||
<StatusIcon className="size-3.5" />
|
||||
<div
|
||||
className={cn(
|
||||
"flex items-center gap-2 rounded-full border border-primary/20 bg-primary/5 px-3 py-1.5 text-xs text-primary",
|
||||
// The spinner is its own motion cue; pulsing too just fades it mid-spin.
|
||||
!isNudging && "animate-pulse",
|
||||
)}
|
||||
>
|
||||
{isNudging ? (
|
||||
// label, not the default "Loading": the spinner is the badge's only
|
||||
// role="status" region, so its name is what gets announced.
|
||||
<Spinner className="size-3.5" label={toolStatus} />
|
||||
) : (
|
||||
<StatusIcon className="size-3.5" />
|
||||
)}
|
||||
<span>{toolStatus}</span>
|
||||
<span className="tabular-nums opacity-60">{elapsed}s</span>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -1604,6 +1604,7 @@ async function autoLoadSmallestModel(): Promise<{
|
|||
? {
|
||||
gpu_ids: effectiveGpuIds ?? undefined,
|
||||
gpu_memory_mode: effectiveGpuMemoryMode,
|
||||
n_parallel: config.nParallel ?? null,
|
||||
}
|
||||
: {}),
|
||||
}))
|
||||
|
|
@ -1637,6 +1638,8 @@ async function autoLoadSmallestModel(): Promise<{
|
|||
gpu_layers: effectiveGpuLayers,
|
||||
n_cpu_moe: effectiveNCpuMoe,
|
||||
gpu_ids: effectiveGpuIds ?? undefined,
|
||||
// Per-model too, or the auto-load reverts a remembered override.
|
||||
n_parallel: config.nParallel ?? null,
|
||||
}
|
||||
: {}),
|
||||
});
|
||||
|
|
@ -1689,6 +1692,11 @@ async function autoLoadSmallestModel(): Promise<{
|
|||
effectiveGpuLayers,
|
||||
config.customContextLength ?? null,
|
||||
);
|
||||
// Slots this auto-load committed. Diffusion ignores --parallel, so a count
|
||||
// there would mint a phantom override a saved preset carries onto a GGUF.
|
||||
const committedSlots = (loadResp.is_diffusion ?? false)
|
||||
? null
|
||||
: (config.nParallel ?? null);
|
||||
useChatRuntimeStore.setState({
|
||||
ggufContextLength: loadResp.context_length ?? 131072,
|
||||
ggufMaxContextLength:
|
||||
|
|
@ -1703,6 +1711,9 @@ async function autoLoadSmallestModel(): Promise<{
|
|||
...resolveToolsEnabledOnLoad(loadResp.supports_tools ?? false),
|
||||
kvCacheDtype: loadResp.cache_type_kv ?? null,
|
||||
loadedKvCacheDtype: loadResp.cache_type_kv ?? null,
|
||||
// Click-time value, not the resolved backend echo (see performLoad).
|
||||
nParallel: committedSlots,
|
||||
loadedNParallel: committedSlots,
|
||||
tensorParallel: loadResp.tensor_parallel ?? false,
|
||||
loadedTensorParallel: loadResp.tensor_parallel ?? false,
|
||||
...loadedGpuMemoryFields(loadResp),
|
||||
|
|
@ -1728,6 +1739,10 @@ async function autoLoadSmallestModel(): Promise<{
|
|||
...resolveToolsEnabledOnLoad(loadResp.supports_tools ?? false),
|
||||
kvCacheDtype: loadResp.cache_type_kv ?? null,
|
||||
loadedKvCacheDtype: loadResp.cache_type_kv ?? null,
|
||||
// GGUF-only and never sent here: a staged override would be saved for
|
||||
// a model that cannot use it.
|
||||
nParallel: null,
|
||||
loadedNParallel: null,
|
||||
tensorParallel: loadResp.tensor_parallel ?? false,
|
||||
loadedTensorParallel: loadResp.tensor_parallel ?? false,
|
||||
// Non-GGUF response: clears any stale GPU baseline a prior manual-GPU
|
||||
|
|
@ -2001,6 +2016,10 @@ async function autoLoadSmallestModel(): Promise<{
|
|||
...resolveToolsEnabledOnLoad(loadResp.supports_tools ?? false),
|
||||
kvCacheDtype: loadResp.cache_type_kv ?? null,
|
||||
loadedKvCacheDtype: loadResp.cache_type_kv ?? null,
|
||||
// The request above omits n_parallel: a staged override left from a
|
||||
// preset would read as applied and be re-sent by the next Apply.
|
||||
nParallel: null,
|
||||
loadedNParallel: null,
|
||||
tensorParallel: loadResp.tensor_parallel ?? false,
|
||||
loadedTensorParallel: loadResp.tensor_parallel ?? false,
|
||||
...loadedGpuMemoryFields(loadResp),
|
||||
|
|
|
|||
|
|
@ -199,6 +199,8 @@ export async function validateModel(
|
|||
// --fit, while a pinned layer count is owned by the user. Tell validate
|
||||
// so it applies the same training-guard policy as /load.
|
||||
gpu_memory_mode: payload.gpu_memory_mode,
|
||||
// Slots scale the KV estimate; keep validate sized like the load.
|
||||
n_parallel: payload.n_parallel,
|
||||
}),
|
||||
});
|
||||
return parseJsonOrThrow<ValidateModelResponse>(response);
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ import { Streamdown } from "streamdown";
|
|||
import { ArtifactHtmlFrame, type ArtifactViewMode } from "./html-frame";
|
||||
import { useChatArtifactsStore } from "./store";
|
||||
import type { ChatArtifact } from "./types";
|
||||
import { getArtifactFilename } from "./types";
|
||||
import { buildArtifactSourceKey, getArtifactFilename } from "./types";
|
||||
|
||||
const COPY_RESET_MS = 2000;
|
||||
const artifactSourceCodePlugin = createCodePlugin({
|
||||
|
|
@ -338,6 +338,8 @@ export function ArtifactSurface({
|
|||
) : (
|
||||
<div className="h-full overflow-auto text-xs leading-relaxed [&_[data-streamdown=code-block]]:!my-0 [&_[data-streamdown=code-block]]:!gap-0 [&_[data-streamdown=code-block]]:!rounded-none [&_[data-streamdown=code-block]]:!border-0 [&_[data-streamdown=code-block]]:!bg-transparent [&_[data-streamdown=code-block]]:!p-0 [&_[data-streamdown=code-block-body]]:!border-0 [&_[data-streamdown=code-block-body]]:!bg-transparent [&_[data-streamdown=code-block-body]]:!p-0 [&_pre]:!m-0 [&_pre]:!bg-transparent [&_pre]:!p-0 [&_pre]:text-xs [&_pre]:leading-relaxed [&_code]:text-xs">
|
||||
<Streamdown
|
||||
// Only computed when the source view is actually on screen.
|
||||
key={buildArtifactSourceKey(artifact)}
|
||||
mode="streaming"
|
||||
plugins={{ code: artifactSourceCodePlugin }}
|
||||
controls={{ code: false }}
|
||||
|
|
|
|||
|
|
@ -41,6 +41,15 @@ export function hashArtifactCode(code: string): string {
|
|||
return (hash >>> 0).toString(36);
|
||||
}
|
||||
|
||||
// The canvas source view keys its Streamdown on this. Streamdown memoizes a code
|
||||
// fence on its node's line/column span, ignoring the text, so equal-line-count
|
||||
// canvases keep the old source. Tool artifact IDs omit the code, so hash it in.
|
||||
export function buildArtifactSourceKey(
|
||||
artifact: Pick<ChatArtifact, "id" | "code">,
|
||||
): string {
|
||||
return `${artifact.id}:${hashArtifactCode(artifact.code)}`;
|
||||
}
|
||||
|
||||
export function createArtifactId(input: ChatArtifactInput): string {
|
||||
const threadSegment = input.threadId || "no-thread";
|
||||
const messageSegment = input.sourceMessageId || "transient";
|
||||
|
|
|
|||
|
|
@ -781,6 +781,7 @@ function GeneralCompareHeader({
|
|||
// Controlled so the body-portaled popover can't linger over another tab off-route.
|
||||
const active = useChatActive();
|
||||
const [selectorOpen, setSelectorOpen] = useState(false);
|
||||
|
||||
const { pinned } = useSidebar();
|
||||
return (
|
||||
<div
|
||||
|
|
@ -3192,7 +3193,7 @@ export function ChatPage({
|
|||
// Provides `active` to ChatRuntimeProvider (drops the message views/composers
|
||||
// while off-route, keeping the runtime alive) and to the compare chrome.
|
||||
<ChatActiveContext.Provider value={active}>
|
||||
<div className="flex min-h-0 min-w-0 flex-1 basis-0 bg-background overflow-hidden">
|
||||
<div className="flex min-h-0 min-w-0 flex-1 basis-0 overflow-hidden bg-background">
|
||||
{/* Portaled surfaces render to document.body, escaping the parent's hidden
|
||||
wrapper, so gate them on `active` to keep them off other tabs. */}
|
||||
{active && <GuidedTour {...tour.tourProps} />}
|
||||
|
|
|
|||
|
|
@ -397,6 +397,7 @@ export function ChatSettingsPanel({
|
|||
const nCpuMoe = useChatRuntimeStore((s) => s.nCpuMoe);
|
||||
const tensorParallel = useChatRuntimeStore((s) => s.tensorParallel);
|
||||
const specDraftNMax = useChatRuntimeStore((s) => s.specDraftNMax);
|
||||
const nParallel = useChatRuntimeStore((s) => s.nParallel);
|
||||
const speculativeType = useChatRuntimeStore((s) => s.speculativeType);
|
||||
const specFallbackReason = useChatRuntimeStore((s) => s.specFallbackReason);
|
||||
const mtpUpdatable =
|
||||
|
|
@ -504,6 +505,7 @@ export function ChatSettingsPanel({
|
|||
tensorParallel,
|
||||
speculativeType,
|
||||
specDraftNMax,
|
||||
nParallel,
|
||||
params.maxSeqLength,
|
||||
]);
|
||||
const activePresetLoadSummary = useMemo(
|
||||
|
|
@ -522,6 +524,7 @@ export function ChatSettingsPanel({
|
|||
tensorParallel,
|
||||
speculativeType,
|
||||
specDraftNMax,
|
||||
nParallel,
|
||||
params.maxSeqLength,
|
||||
],
|
||||
);
|
||||
|
|
|
|||
|
|
@ -567,6 +567,8 @@ export function useChatModelRuntime() {
|
|||
applyActiveModelStatusToStore(residentStatus, {
|
||||
previousCheckpoint: selectedCheckpoint,
|
||||
previousGgufVariant,
|
||||
// Id and variant matched above: same model, only the tab moved.
|
||||
readoptingSameModel: true,
|
||||
});
|
||||
syncModelCapabilities(modelId, residentStatus);
|
||||
return;
|
||||
|
|
@ -669,6 +671,14 @@ export function useChatModelRuntime() {
|
|||
let previousWasUnloaded = false;
|
||||
const pendingLoadConfig =
|
||||
typeof selection !== "string" ? selection.config : undefined;
|
||||
// The outgoing model's slot INTENT (blank = follow the server
|
||||
// default), which the resolved baseline cannot express. previousConfig
|
||||
// is the snapshot the picker took before pre-applying the target's
|
||||
// config, so the live control is only the outgoing one without it.
|
||||
const previousNParallel =
|
||||
typeof selection !== "string" && selection.previousConfig
|
||||
? (selection.previousConfig.nParallel ?? null)
|
||||
: useChatRuntimeStore.getState().nParallel;
|
||||
if (pendingLoadConfig) {
|
||||
applyPerModelConfigToRuntime(pendingLoadConfig);
|
||||
}
|
||||
|
|
@ -761,6 +771,8 @@ export function useChatModelRuntime() {
|
|||
: stateBeforeUnload.speculativeType;
|
||||
let loadSpecDraftNMax =
|
||||
pendingLoadConfig?.specDraftNMax ?? stateBeforeUnload.specDraftNMax;
|
||||
let loadNParallel =
|
||||
pendingLoadConfig?.nParallel ?? stateBeforeUnload.nParallel;
|
||||
try {
|
||||
// Lightweight pre-flight validation: avoid unloading a working model
|
||||
// if the new identifier is clearly invalid (e.g. bad HF id / path).
|
||||
|
|
@ -792,6 +804,10 @@ export function useChatModelRuntime() {
|
|||
const validateGpuLayers = resetsPerModelSettings
|
||||
? GPU_LAYERS_AUTO
|
||||
: loadGpuLayers;
|
||||
// Per-model: the reset re-baselines to the staged config, like the load.
|
||||
const validateNParallel = resetsPerModelSettings
|
||||
? (pendingLoadConfig?.nParallel ?? null)
|
||||
: loadNParallel;
|
||||
const validateMaxSeqLength = resolveFitMaxSeqLength(
|
||||
isGguf,
|
||||
loadGpuMemoryMode,
|
||||
|
|
@ -820,7 +836,12 @@ export function useChatModelRuntime() {
|
|||
cache_type_kv: loadKvCacheDtype,
|
||||
tensor_parallel: loadTensorParallel,
|
||||
gpu_ids: validateGpuIds ?? undefined,
|
||||
...(isGguf ? { gpu_memory_mode: loadGpuMemoryMode } : {}),
|
||||
...(isGguf
|
||||
? {
|
||||
gpu_memory_mode: loadGpuMemoryMode,
|
||||
n_parallel: validateNParallel,
|
||||
}
|
||||
: {}),
|
||||
});
|
||||
// Upgrade consent runs before the security dialogs; Accept installs and the load continues.
|
||||
if (validation.requires_transformers_upgrade) {
|
||||
|
|
@ -903,6 +924,10 @@ export function useChatModelRuntime() {
|
|||
loadedSpeculativeType: persistedSpeculativeType,
|
||||
specDraftNMax: null,
|
||||
loadedSpecDraftNMax: null,
|
||||
// Per-model too: a different model follows the server default
|
||||
// unless its staged config overrides it.
|
||||
nParallel: null,
|
||||
loadedNParallel: null,
|
||||
// Per-model GPU knobs must not follow onto a different model
|
||||
// (gpuMemoryMode is a standing preference and is kept).
|
||||
selectedGpuIds: null,
|
||||
|
|
@ -918,6 +943,7 @@ export function useChatModelRuntime() {
|
|||
? normalizeSpeculativeType(pendingLoadConfig.speculativeType)
|
||||
: persistedSpeculativeType;
|
||||
loadSpecDraftNMax = pendingLoadConfig?.specDraftNMax ?? null;
|
||||
loadNParallel = pendingLoadConfig?.nParallel ?? null;
|
||||
// Keep the click-time snapshot in lock-step with the store reset so
|
||||
// the load below sizes against the cleared per-model knobs, not the
|
||||
// previous model's (gpuMemoryMode is standing, so left as captured).
|
||||
|
|
@ -984,6 +1010,8 @@ export function useChatModelRuntime() {
|
|||
cache_type_kv: loadKvCacheDtype,
|
||||
speculative_type: loadSpeculativeType,
|
||||
spec_draft_n_max: loadSpecDraftNMax,
|
||||
// GGUF-only: slots mean nothing for a transformers load.
|
||||
n_parallel: isGguf ? loadNParallel : null,
|
||||
tensor_parallel: loadTensorParallel,
|
||||
gpu_memory_mode: loadGpuMemoryMode,
|
||||
gpu_layers: loadGpuLayers,
|
||||
|
|
@ -1034,6 +1062,14 @@ export function useChatModelRuntime() {
|
|||
const loadedSpec = normalizeSpeculativeType(
|
||||
loadResponse.speculative_type,
|
||||
);
|
||||
// Slots the load actually committed. Non-GGUF never sends them and
|
||||
// diffusion ignores --parallel, so a click-time count on either
|
||||
// would mint a phantom override a saved preset carries onto a GGUF.
|
||||
const committedSlots =
|
||||
(loadResponse.is_gguf ?? false) &&
|
||||
!(loadResponse.is_diffusion ?? false)
|
||||
? (loadNParallel ?? null)
|
||||
: null;
|
||||
const nativeCtx = loadResponse.is_gguf
|
||||
? (loadResponse.context_length ?? 131072)
|
||||
: null;
|
||||
|
|
@ -1109,6 +1145,10 @@ export function useChatModelRuntime() {
|
|||
loadedSpeculativeType: loadedSpec,
|
||||
specDraftNMax: loadResponse.spec_draft_n_max ?? null,
|
||||
loadedSpecDraftNMax: loadResponse.spec_draft_n_max ?? null,
|
||||
// Keep the click-time value: the echo is the resolved count, and
|
||||
// adopting it would pin a blank "server default" control.
|
||||
nParallel: committedSlots,
|
||||
loadedNParallel: committedSlots,
|
||||
customContextLength: keepCustomCtx,
|
||||
loadedCustomContextLength: keepCustomCtx,
|
||||
defaultChatTemplate: loadResponse.chat_template ?? null,
|
||||
|
|
@ -1211,6 +1251,7 @@ export function useChatModelRuntime() {
|
|||
stateBeforeUnload.loadedSpeculativeType,
|
||||
spec_draft_n_max:
|
||||
stateBeforeUnload.loadedSpecDraftNMax,
|
||||
n_parallel: stateBeforeUnload.loadedNParallel,
|
||||
// Restore the previous model in the split mode it was running,
|
||||
// not the default layer split.
|
||||
tensor_parallel: stateBeforeUnload.loadedTensorParallel ?? false,
|
||||
|
|
@ -1237,6 +1278,9 @@ export function useChatModelRuntime() {
|
|||
// model's; the loaded baselines below come from its reload echo.
|
||||
speculativeType: stateBeforeUnload.loadedSpeculativeType ?? null,
|
||||
specDraftNMax: stateBeforeUnload.loadedSpecDraftNMax ?? null,
|
||||
// Control keeps its intent; only the baseline takes the echo.
|
||||
nParallel: previousNParallel,
|
||||
loadedNParallel: stateBeforeUnload.loadedNParallel ?? null,
|
||||
loadedSpeculativeType: rollbackSpeculativeType,
|
||||
loadedSpecDraftNMax:
|
||||
rollbackResponse.spec_draft_n_max ?? null,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
// Barrel import (lint rule); the model-picker cycle is fine because the call
|
||||
// happens at runtime, not module eval.
|
||||
import { resolveInitialConfig } from "@/features/model-picker";
|
||||
import { getInferenceStatus } from "../api/chat-api";
|
||||
import {
|
||||
mergeBackendRecommendedInference,
|
||||
|
|
@ -131,6 +134,9 @@ export type ApplyInferenceStatusOptions = {
|
|||
* status -- without it a variant-only switch underneath the tab reads as
|
||||
* steady state and the hydration reseed keeps the old quant's baselines. */
|
||||
previousGgufVariant?: string | null;
|
||||
/** The caller verified the status is the model this tab just picked, so the
|
||||
* slot control it holds belongs to that model and must survive. */
|
||||
readoptingSameModel?: boolean;
|
||||
};
|
||||
|
||||
/** Mirror refresh() hydration so adopted CLI models get reasoning/tools flags. */
|
||||
|
|
@ -201,6 +207,22 @@ export function applyActiveModelStatusToStore(
|
|||
// While a load is in flight, performLoad owns the load params. Seeding them
|
||||
// from a stale poll here would clobber the values the load dialog just set.
|
||||
const seedLoadParams = !prevState.modelLoading;
|
||||
// A model/variant change underneath this tab, as opposed to re-adopting the
|
||||
// model the tab just picked, where hydratingExistingModel fires on the stale
|
||||
// checkpoint. The echo cannot stand in: a new model can report the old count.
|
||||
const slotsModelChanged =
|
||||
hydratingExistingModel && !options.readoptingSameModel;
|
||||
// This model's remembered override, read only on a fresh store or a model
|
||||
// change, so a steady poll cannot re-pin a control the user just blanked.
|
||||
const slotsUnseeded =
|
||||
prevState.loadedNParallel === null && prevState.nParallel === null;
|
||||
const remembered =
|
||||
status.is_gguf && (slotsUnseeded || slotsModelChanged)
|
||||
? resolveInitialConfig(checkpointId, status.gguf_variant ?? null)
|
||||
: null;
|
||||
const rememberedNParallel = remembered?.remembered
|
||||
? (remembered.config.nParallel ?? null)
|
||||
: null;
|
||||
// A Manual + Auto-layers load sent its positive context pin as max_seq_length,
|
||||
// and status only exposes the RESOLVED context; re-seed the pin from the
|
||||
// requested value (parity with the load paths' keepCustomCtx). Baselines
|
||||
|
|
@ -322,6 +344,35 @@ export function applyActiveModelStatusToStore(
|
|||
tensorParallel: status.tensor_parallel,
|
||||
loadedTensorParallel: status.tensor_parallel,
|
||||
}),
|
||||
// Baseline only, never the control: the echo is the RESOLVED count and would
|
||||
// pin a blank "server default" control. The rollback re-sends the baseline,
|
||||
// so without this a rollback after a tab reload loses the override.
|
||||
...(seedLoadParams &&
|
||||
status.requested_parallel_slots != null &&
|
||||
(prevState.loadedNParallel === null || hydratingExistingModel) && {
|
||||
loadedNParallel: status.requested_parallel_slots,
|
||||
}),
|
||||
// A slotless model must not keep the previous GGUF's baseline: the rollback
|
||||
// re-sends it. /status omits the echo for non-GGUF and sends an explicit
|
||||
// null for diffusion, so an absent field on a GGUF is an older backend.
|
||||
...(seedLoadParams &&
|
||||
(status.is_gguf === false || status.requested_parallel_slots === null) && {
|
||||
loadedNParallel: null,
|
||||
}),
|
||||
// Per-model: a change underneath this tab blanks the control like
|
||||
// performLoad's cross-model reset, or the old count follows onto the new
|
||||
// model. The baseline above still carries the rollback.
|
||||
...(seedLoadParams && slotsModelChanged && { nParallel: null }),
|
||||
// AFTER that clear, which both a first hydration and a model change trip:
|
||||
// either would leave the control blank while the model runs on a remembered
|
||||
// override, so the next Apply would save the blank over it. Adopted only
|
||||
// when the running count matches, proving it is this model's own.
|
||||
...(seedLoadParams &&
|
||||
(slotsUnseeded || slotsModelChanged) &&
|
||||
rememberedNParallel != null &&
|
||||
rememberedNParallel === status.requested_parallel_slots && {
|
||||
nParallel: rememberedNParallel,
|
||||
}),
|
||||
// Re-seed on first hydration, model/variant changes, or a same-model backend
|
||||
// placement change. gpuStatusFields preserves dirty local edits in the last
|
||||
// case while advancing their loaded baselines.
|
||||
|
|
|
|||
|
|
@ -12,6 +12,8 @@ import {
|
|||
DEFAULT_MAX_SEQ_LENGTH,
|
||||
KV_CACHE_DTYPES,
|
||||
MTP_SPECULATIVE_TYPES,
|
||||
N_PARALLEL_MAX,
|
||||
N_PARALLEL_MIN,
|
||||
SPECULATIVE_TYPES,
|
||||
normalizeMaxSeqLength,
|
||||
type PerModelConfig,
|
||||
|
|
@ -30,6 +32,7 @@ export type PresetLoadConfig = Pick<
|
|||
| "kvCacheDtype"
|
||||
| "speculativeType"
|
||||
| "specDraftNMax"
|
||||
| "nParallel"
|
||||
| "tensorParallel"
|
||||
| "gpuMemoryMode"
|
||||
| "gpuLayers"
|
||||
|
|
@ -45,6 +48,7 @@ export const EMPTY_PRESET_LOAD_CONFIG: PresetLoadConfig = {
|
|||
kvCacheDtype: null,
|
||||
speculativeType: null,
|
||||
specDraftNMax: null,
|
||||
nParallel: null,
|
||||
tensorParallel: false,
|
||||
};
|
||||
|
||||
|
|
@ -107,6 +111,14 @@ export function normalizePresetLoadConfig(
|
|||
? speculativeType
|
||||
: null,
|
||||
specDraftNMax,
|
||||
nParallel:
|
||||
typeof partial.nParallel === "number" &&
|
||||
Number.isFinite(partial.nParallel)
|
||||
? Math.max(
|
||||
N_PARALLEL_MIN,
|
||||
Math.min(N_PARALLEL_MAX, Math.round(partial.nParallel)),
|
||||
)
|
||||
: null,
|
||||
tensorParallel:
|
||||
typeof partial.tensorParallel === "boolean"
|
||||
? partial.tensorParallel
|
||||
|
|
@ -151,6 +163,7 @@ export function capturePresetLoadConfig(): PresetLoadConfig | undefined {
|
|||
kvCacheDtype: snapshot.kvCacheDtype ?? null,
|
||||
speculativeType: normalizeSpeculativeType(snapshot.speculativeType),
|
||||
specDraftNMax: snapshot.specDraftNMax ?? null,
|
||||
nParallel: snapshot.nParallel ?? null,
|
||||
tensorParallel: snapshot.tensorParallel ?? false,
|
||||
...(snapshot.gpuMemoryMode === "manual"
|
||||
? { gpuMemoryMode: "manual" as const }
|
||||
|
|
@ -206,6 +219,7 @@ export function applyPresetLoadConfig(
|
|||
kvCacheDtype: config.kvCacheDtype ?? null,
|
||||
speculativeType: config.speculativeType ?? null,
|
||||
specDraftNMax: config.specDraftNMax ?? null,
|
||||
nParallel: config.nParallel ?? null,
|
||||
tensorParallel: config.tensorParallel ?? false,
|
||||
chatTemplateOverride: null,
|
||||
gpuMemoryMode: config.gpuMemoryMode,
|
||||
|
|
@ -231,6 +245,9 @@ export function formatPresetLoadConfigSummary(
|
|||
if (config.speculativeType && config.speculativeType !== "auto") {
|
||||
parts.push(`Spec ${config.speculativeType}`);
|
||||
}
|
||||
if (config.nParallel != null) {
|
||||
parts.push(`${config.nParallel} slots`);
|
||||
}
|
||||
if (config.gpuMemoryMode === "manual") {
|
||||
parts.push("GPU manual");
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1130,6 +1130,8 @@ export function SharedComposer({
|
|||
? {
|
||||
gpu_ids: effectiveSelectedGpuIds ?? undefined,
|
||||
gpu_memory_mode: effectiveGpuMemoryMode,
|
||||
// Slots scale the KV estimate; keep validate sized like the load.
|
||||
n_parallel: ownConfig.nParallel ?? null,
|
||||
}
|
||||
: {}),
|
||||
});
|
||||
|
|
@ -1198,6 +1200,7 @@ export function SharedComposer({
|
|||
n_cpu_moe: effectiveNCpuMoe,
|
||||
tensor_split: compareLoadKnobs.splitRatio ?? undefined,
|
||||
gpu_ids: effectiveSelectedGpuIds ?? undefined,
|
||||
n_parallel: ownConfig.nParallel ?? null,
|
||||
}
|
||||
: {}),
|
||||
});
|
||||
|
|
@ -1229,6 +1232,12 @@ export function SharedComposer({
|
|||
effectiveCustomContextLength,
|
||||
)
|
||||
: null;
|
||||
// Slots this compare load committed. Diffusion ignores --parallel, so a
|
||||
// count there would mint a phantom override a preset carries onto a GGUF.
|
||||
const committedSlots =
|
||||
targetIsGguf && !(resp.is_diffusion ?? false)
|
||||
? (ownConfig.nParallel ?? null)
|
||||
: null;
|
||||
useChatRuntimeStore.setState({
|
||||
supportsReasoning: resp.supports_reasoning ?? false,
|
||||
reasoningAlwaysOn: resp.reasoning_always_on ?? false,
|
||||
|
|
@ -1237,6 +1246,9 @@ export function SharedComposer({
|
|||
supportsTools: resp.supports_tools ?? false,
|
||||
kvCacheDtype: resp.cache_type_kv ?? null,
|
||||
loadedKvCacheDtype: resp.cache_type_kv ?? null,
|
||||
// Click-time value, not the resolved echo (see the single-model load).
|
||||
nParallel: committedSlots,
|
||||
loadedNParallel: committedSlots,
|
||||
tensorParallel: resp.tensor_parallel ?? false,
|
||||
loadedTensorParallel: resp.tensor_parallel ?? false,
|
||||
defaultChatTemplate: resp.chat_template ?? null,
|
||||
|
|
|
|||
|
|
@ -968,6 +968,12 @@ type ChatRuntimeStore = {
|
|||
/** User --spec-draft-n-max override (null = platform default). */
|
||||
specDraftNMax: number | null;
|
||||
loadedSpecDraftNMax: number | null;
|
||||
/** User --parallel slots override for GGUF loads (null = server default).
|
||||
* Never re-seeded from an echo: the resolved count would pin a blank control. */
|
||||
nParallel: number | null;
|
||||
/** Slots the last successful load sent (null = default); the rollback
|
||||
* re-sends it so a failed switch can't lose the override. */
|
||||
loadedNParallel: number | null;
|
||||
/** Tensor-parallel split (--split-mode tensor) toggle, GGUF multi-GPU only. */
|
||||
tensorParallel: boolean;
|
||||
/** Backend-reported tensor-parallel state; null until first hydrated. */
|
||||
|
|
@ -1491,6 +1497,8 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
|
|||
specFallbackReason: null,
|
||||
specDraftNMax: null,
|
||||
loadedSpecDraftNMax: null,
|
||||
nParallel: null,
|
||||
loadedNParallel: null,
|
||||
tensorParallel: false,
|
||||
loadedTensorParallel: null,
|
||||
gpuMemoryMode: readPersistedGpuMemoryMode(),
|
||||
|
|
@ -1874,6 +1882,8 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
|
|||
specFallbackReason: null,
|
||||
specDraftNMax: null,
|
||||
loadedSpecDraftNMax: null,
|
||||
nParallel: null,
|
||||
loadedNParallel: null,
|
||||
tensorParallel: false,
|
||||
loadedTensorParallel: null,
|
||||
// Standing preference: survives unload, unlike the per-model knobs above.
|
||||
|
|
|
|||
|
|
@ -65,6 +65,11 @@ export interface LoadModelRequest {
|
|||
* when speculative_type resolves to "mtp" or "mtp+ngram".
|
||||
*/
|
||||
spec_draft_n_max?: number | null;
|
||||
/**
|
||||
* Parallel decode slots for llama-server (--parallel), 1..64. Omit/null =
|
||||
* the launch default. The VRAM fitter may launch fewer to stay on GPU.
|
||||
*/
|
||||
n_parallel?: number | null;
|
||||
/**
|
||||
* Split the model across GPUs by tensor (--split-mode tensor) instead
|
||||
* of by layer for GGUF models. Multi-GPU only; no effect on a single GPU.
|
||||
|
|
@ -202,6 +207,12 @@ export interface LoadModelResponse {
|
|||
gpu_ids?: number[] | null;
|
||||
/** User-requested GPU placement pool before fit-time narrowing. */
|
||||
requested_gpu_ids?: number[] | null;
|
||||
/** Slots the load was invoked with (else the --parallel default). Null for
|
||||
* non-GGUF loads. */
|
||||
requested_parallel_slots?: number | null;
|
||||
/** Slots llama-server actually runs, after any fit-time reduction. Null for
|
||||
* non-GGUF loads. */
|
||||
parallel_slots?: number | null;
|
||||
}
|
||||
|
||||
export interface UnloadModelRequest {
|
||||
|
|
@ -263,6 +274,12 @@ export interface InferenceStatusResponse {
|
|||
gpu_ids?: number[] | null;
|
||||
/** User-requested GPU placement pool before fit-time narrowing. */
|
||||
requested_gpu_ids?: number[] | null;
|
||||
/** Slots the active load was invoked with (else the --parallel default).
|
||||
* Null when no GGUF model is loaded. */
|
||||
requested_parallel_slots?: number | null;
|
||||
/** Slots llama-server actually runs, after any fit-time reduction. Null when
|
||||
* no GGUF model is loaded. */
|
||||
parallel_slots?: number | null;
|
||||
n_layers?: number | null;
|
||||
/** Model's MoE expert-layer count (the n_cpu_moe ceiling); 0 if not MoE. */
|
||||
n_moe_layers?: number;
|
||||
|
|
|
|||
15
studio/frontend/src/features/chat/utils/tool-status.ts
Normal file
15
studio/frontend/src/features/chat/utils/tool-status.ts
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
/** Mirrors NUDGE_TOOL_CALLS_STATUS in backend core/inference/tool_call_parser.py; keep in sync. */
|
||||
export const NUDGE_TOOL_CALLS_STATUS = "Nudging tool calls";
|
||||
|
||||
export type ToolStatusKind = "nudge" | "terminal" | "web";
|
||||
|
||||
/** Which glyph the badge shows: exact match for the nudge, "Running" prefix for sandbox tools, globe otherwise. */
|
||||
export function toolStatusKind(status: string): ToolStatusKind {
|
||||
if (status === NUDGE_TOOL_CALLS_STATUS) {
|
||||
return "nudge";
|
||||
}
|
||||
return status.startsWith("Running") ? "terminal" : "web";
|
||||
}
|
||||
|
|
@ -0,0 +1,92 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
import { isTauri } from "@/lib/api-base";
|
||||
import { useNavigate } from "@tanstack/react-router";
|
||||
import { useEffect } from "react";
|
||||
|
||||
import { createDeepLinkIntentGate } from "./deep-link-intent";
|
||||
import { parseUnslothDeepLink } from "./parse-deep-link";
|
||||
|
||||
const acceptIntent = createDeepLinkIntentGate(2_000);
|
||||
|
||||
async function restoreMainWindow(): Promise<void> {
|
||||
const { getCurrentWindow } = await import("@tauri-apps/api/window");
|
||||
const window = getCurrentWindow();
|
||||
await window.show();
|
||||
await window.unminimize();
|
||||
await window.setFocus();
|
||||
}
|
||||
|
||||
export function DeepLinkHandler() {
|
||||
const navigate = useNavigate();
|
||||
|
||||
useEffect(() => {
|
||||
if (!isTauri) return;
|
||||
|
||||
let disposed = false;
|
||||
let receivedLiveIntent = false;
|
||||
let unlisten: (() => void) | undefined;
|
||||
|
||||
const handleUrls = (urls: string[]): boolean => {
|
||||
if (disposed) return false;
|
||||
|
||||
let hasValidIntent = false;
|
||||
let intent: ReturnType<typeof parseUnslothDeepLink> = null;
|
||||
|
||||
let intentSequence: number | null = null;
|
||||
for (const rawUrl of urls) {
|
||||
const parsed = parseUnslothDeepLink(rawUrl);
|
||||
if (!parsed) continue;
|
||||
hasValidIntent = true;
|
||||
const sequence = acceptIntent(parsed.model, parsed.file);
|
||||
if (sequence !== null) {
|
||||
intent = parsed;
|
||||
intentSequence = sequence;
|
||||
}
|
||||
}
|
||||
if (!intent || intentSequence === null) return hasValidIntent;
|
||||
|
||||
void restoreMainWindow().catch(() => undefined);
|
||||
void navigate({
|
||||
to: "/hub",
|
||||
search: {
|
||||
tab: "discover",
|
||||
kind: "models",
|
||||
model: intent.model,
|
||||
file: intent.file,
|
||||
|
||||
intent: intentSequence,
|
||||
},
|
||||
});
|
||||
return true;
|
||||
};
|
||||
|
||||
async function subscribe() {
|
||||
const { getCurrent, onOpenUrl } =
|
||||
await import("@tauri-apps/plugin-deep-link");
|
||||
if (disposed) return;
|
||||
|
||||
const cleanup = await onOpenUrl((urls) => {
|
||||
if (handleUrls(urls)) receivedLiveIntent = true;
|
||||
});
|
||||
if (disposed) {
|
||||
cleanup();
|
||||
return;
|
||||
}
|
||||
unlisten = cleanup;
|
||||
|
||||
const currentUrls = await getCurrent();
|
||||
if (currentUrls && !receivedLiveIntent) handleUrls(currentUrls);
|
||||
}
|
||||
|
||||
void subscribe().catch(() => undefined);
|
||||
|
||||
return () => {
|
||||
disposed = true;
|
||||
unlisten?.();
|
||||
};
|
||||
}, [navigate]);
|
||||
|
||||
return null;
|
||||
}
|
||||
24
studio/frontend/src/features/deep-links/deep-link-intent.ts
Normal file
24
studio/frontend/src/features/deep-links/deep-link-intent.ts
Normal file
|
|
@ -0,0 +1,24 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
export function createDeepLinkIntentGate(
|
||||
deduplicationWindowMs: number,
|
||||
now: () => number = Date.now,
|
||||
) {
|
||||
let lastIntent: { key: string; handledAt: number } | null = null;
|
||||
let sequence = 0;
|
||||
|
||||
return (model: string, file?: string): number | null => {
|
||||
const handledAt = now();
|
||||
const key = `${model}\0${file ?? ""}`;
|
||||
if (
|
||||
lastIntent?.key === key &&
|
||||
handledAt - lastIntent.handledAt < deduplicationWindowMs
|
||||
) {
|
||||
return null;
|
||||
}
|
||||
lastIntent = { key, handledAt };
|
||||
sequence += 1;
|
||||
return sequence;
|
||||
};
|
||||
}
|
||||
4
studio/frontend/src/features/deep-links/index.ts
Normal file
4
studio/frontend/src/features/deep-links/index.ts
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
export { DeepLinkHandler } from "./deep-link-handler";
|
||||
101
studio/frontend/src/features/deep-links/parse-deep-link.ts
Normal file
101
studio/frontend/src/features/deep-links/parse-deep-link.ts
Normal file
|
|
@ -0,0 +1,101 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
const MAX_REPO_ID_SEGMENT_LENGTH = 96;
|
||||
const MAX_GGUF_FILE_LENGTH = 512;
|
||||
const REPO_SEGMENT = /^[A-Za-z0-9](?:[A-Za-z0-9._-]*[A-Za-z0-9])?$/;
|
||||
function hasControlCharacters(value: string): boolean {
|
||||
return [...value].some((character) => {
|
||||
const codePoint = character.codePointAt(0) ?? 0;
|
||||
return codePoint <= 0x1f || codePoint === 0x7f;
|
||||
});
|
||||
}
|
||||
|
||||
export interface UnslothDeepLinkIntent {
|
||||
model: string;
|
||||
file?: string;
|
||||
}
|
||||
|
||||
function isValidRepoSegment(segment: string): boolean {
|
||||
return (
|
||||
segment.length <= MAX_REPO_ID_SEGMENT_LENGTH &&
|
||||
REPO_SEGMENT.test(segment) &&
|
||||
!segment.includes("--") &&
|
||||
!segment.includes("..")
|
||||
);
|
||||
}
|
||||
|
||||
function isValidGgufFile(file: string): boolean {
|
||||
if (
|
||||
file.length === 0 ||
|
||||
file.length > MAX_GGUF_FILE_LENGTH ||
|
||||
file !== file.trim() ||
|
||||
hasControlCharacters(file) ||
|
||||
file.includes("\\") ||
|
||||
file.startsWith("/") ||
|
||||
!file.toLowerCase().endsWith(".gguf")
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
return file
|
||||
.split("/")
|
||||
.every((segment) => segment !== "" && segment !== "." && segment !== "..");
|
||||
}
|
||||
|
||||
export function parseUnslothDeepLink(
|
||||
rawUrl: string,
|
||||
): UnslothDeepLinkIntent | null {
|
||||
const queryIndex = rawUrl.indexOf("?");
|
||||
const target = queryIndex === -1 ? rawUrl : rawUrl.slice(0, queryIndex);
|
||||
if (
|
||||
target !== "unsloth://open_from_hf" &&
|
||||
target !== "unsloth://open_from_hf/"
|
||||
) {
|
||||
return null;
|
||||
}
|
||||
|
||||
let url: URL;
|
||||
try {
|
||||
url = new URL(rawUrl);
|
||||
} catch {
|
||||
return null;
|
||||
}
|
||||
|
||||
if (
|
||||
url.protocol !== "unsloth:" ||
|
||||
url.hostname !== "open_from_hf" ||
|
||||
(url.pathname !== "" && url.pathname !== "/") ||
|
||||
url.username !== "" ||
|
||||
url.password !== "" ||
|
||||
url.port !== "" ||
|
||||
url.hash !== ""
|
||||
) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const keys = [...url.searchParams.keys()];
|
||||
if (
|
||||
keys.length < 1 ||
|
||||
keys.length > 2 ||
|
||||
!keys.includes("model") ||
|
||||
new Set(keys).size !== keys.length ||
|
||||
keys.some((key) => key !== "model" && key !== "file")
|
||||
) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const model = url.searchParams.get("model") ?? "";
|
||||
const segments = model.split("/");
|
||||
if (
|
||||
model.endsWith(".git") ||
|
||||
segments.length !== 2 ||
|
||||
!segments.every(isValidRepoSegment)
|
||||
) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const file = url.searchParams.get("file");
|
||||
if (file !== null && !isValidGgufFile(file)) return null;
|
||||
|
||||
return file === null ? { model } : { model, file };
|
||||
}
|
||||
|
|
@ -15,6 +15,9 @@ export function DownloadSection({
|
|||
canRun = true,
|
||||
isActive,
|
||||
activeQuant,
|
||||
preferredGgufFile = null,
|
||||
|
||||
preferredGgufFileIntent = 0,
|
||||
isLoadingThisModel,
|
||||
gpuGb,
|
||||
systemRamGb,
|
||||
|
|
@ -35,6 +38,9 @@ export function DownloadSection({
|
|||
canRun?: boolean;
|
||||
isActive: boolean;
|
||||
activeQuant: string | null;
|
||||
preferredGgufFile?: string | null;
|
||||
|
||||
preferredGgufFileIntent?: number;
|
||||
isLoadingThisModel: boolean;
|
||||
gpuGb?: number;
|
||||
systemRamGb?: number;
|
||||
|
|
@ -46,12 +52,15 @@ export function DownloadSection({
|
|||
onTrain?: () => void;
|
||||
onChange?: () => void;
|
||||
}) {
|
||||
if (isGguf) {
|
||||
if (isGguf || preferredGgufFile) {
|
||||
return (
|
||||
<GgufDownloadCard
|
||||
repoId={repoId}
|
||||
isActive={isActive}
|
||||
activeQuant={activeQuant}
|
||||
preferredFile={preferredGgufFile}
|
||||
|
||||
preferredFileIntent={preferredGgufFileIntent}
|
||||
isLoadingThisModel={isLoadingThisModel}
|
||||
gpuGb={gpuGb}
|
||||
systemRamGb={systemRamGb}
|
||||
|
|
|
|||
|
|
@ -57,6 +57,10 @@ import { useOnlineStatus } from "../hooks/use-online-status";
|
|||
import { type GgufVariantDetail, deleteCachedModel } from "../inventory";
|
||||
import { formatBytes } from "../lib/format";
|
||||
import { type GgufFitClass, classifyGgufFit } from "../lib/gguf-fit";
|
||||
import {
|
||||
ggufFilenamesMatch,
|
||||
ggufSelectionOverrideMatchesIntent,
|
||||
} from "../lib/gguf-filename";
|
||||
import {
|
||||
ggufVariantDisplayLabel,
|
||||
ggufVariantDownloadSizeBytes,
|
||||
|
|
@ -540,6 +544,9 @@ export function GgufDownloadCard({
|
|||
repoId,
|
||||
isActive,
|
||||
activeQuant,
|
||||
preferredFile = null,
|
||||
|
||||
preferredFileIntent = 0,
|
||||
isLoadingThisModel,
|
||||
gpuGb,
|
||||
systemRamGb,
|
||||
|
|
@ -553,6 +560,9 @@ export function GgufDownloadCard({
|
|||
repoId: string;
|
||||
isActive: boolean;
|
||||
activeQuant: string | null;
|
||||
preferredFile?: string | null;
|
||||
|
||||
preferredFileIntent?: number;
|
||||
isLoadingThisModel: boolean;
|
||||
gpuGb?: number;
|
||||
systemRamGb?: number;
|
||||
|
|
@ -579,9 +589,25 @@ export function GgufDownloadCard({
|
|||
repoId: string;
|
||||
quant: string | null;
|
||||
userPicked?: boolean;
|
||||
preferredFile?: string | null;
|
||||
|
||||
preferredFileIntent?: number;
|
||||
}>(() => ({ repoId, quant: null }));
|
||||
const preferredQuant = preferredFile
|
||||
? (variants?.find((variant) =>
|
||||
ggufFilenamesMatch(variant.filename, preferredFile),
|
||||
)?.quant ?? null)
|
||||
: null;
|
||||
const selectedQuantOverride =
|
||||
selectedQuantState.repoId === repoId ? selectedQuantState.quant : null;
|
||||
selectedQuantState.repoId === repoId &&
|
||||
ggufSelectionOverrideMatchesIntent(
|
||||
preferredFile,
|
||||
preferredFileIntent,
|
||||
selectedQuantState.preferredFile,
|
||||
selectedQuantState.preferredFileIntent,
|
||||
)
|
||||
? selectedQuantState.quant
|
||||
: preferredQuant;
|
||||
const [open, setOpen] = useState(false);
|
||||
const [deleteTarget, setDeleteTarget] = useState<string | null>(null);
|
||||
const [updateTarget, setUpdateTarget] = useState<string | null>(null);
|
||||
|
|
@ -732,10 +758,13 @@ export function GgufDownloadCard({
|
|||
repoId,
|
||||
quant,
|
||||
userPicked: true,
|
||||
preferredFile,
|
||||
|
||||
preferredFileIntent,
|
||||
});
|
||||
setOpen(false);
|
||||
},
|
||||
[repoId],
|
||||
[preferredFile, preferredFileIntent, repoId],
|
||||
);
|
||||
const handleDeleteVariant = useCallback((quant: string) => {
|
||||
setDeleteTarget(quant);
|
||||
|
|
|
|||
|
|
@ -38,6 +38,11 @@ import {
|
|||
deleteCachedModel,
|
||||
} from "../inventory";
|
||||
import { formatBytes } from "../lib/format";
|
||||
|
||||
import {
|
||||
ggufFilenamesMatch,
|
||||
ggufSelectionOverrideMatchesIntent,
|
||||
} from "../lib/gguf-filename";
|
||||
import {
|
||||
ggufVariantDisplayLabel,
|
||||
sortLocalGgufVariants,
|
||||
|
|
@ -88,6 +93,9 @@ interface LocalOnDeviceCardProps {
|
|||
activeGgufVariant?: string | null;
|
||||
isLoading: boolean;
|
||||
loadingPhase?: "downloading" | "starting";
|
||||
preferredFile?: string | null;
|
||||
preferredFileIntent?: number;
|
||||
|
||||
gpuGb?: number;
|
||||
systemRamGb?: number;
|
||||
unsupportedReason?: string | null;
|
||||
|
|
@ -210,6 +218,9 @@ export function LocalOnDeviceCard({
|
|||
activeGgufVariant = null,
|
||||
isLoading,
|
||||
loadingPhase,
|
||||
preferredFile = null,
|
||||
preferredFileIntent = 0,
|
||||
|
||||
gpuGb,
|
||||
systemRamGb,
|
||||
unsupportedReason,
|
||||
|
|
@ -285,6 +296,8 @@ export function LocalOnDeviceCard({
|
|||
const [selectedVariantState, setSelectedVariantState] = useState<{
|
||||
key: string;
|
||||
quant: string | null;
|
||||
preferredFile?: string | null;
|
||||
preferredFileIntent?: number;
|
||||
}>(() => ({
|
||||
key: variantKey,
|
||||
quant: null,
|
||||
|
|
@ -328,8 +341,21 @@ export function LocalOnDeviceCard({
|
|||
systemRamGb,
|
||||
],
|
||||
);
|
||||
const preferredQuant = preferredFile
|
||||
? (variants?.find((variant) =>
|
||||
ggufFilenamesMatch(variant.filename, preferredFile),
|
||||
)?.quant ?? null)
|
||||
: null;
|
||||
const selectedVariantOverride =
|
||||
selectedVariantState.key === variantKey ? selectedVariantState.quant : null;
|
||||
selectedVariantState.key === variantKey &&
|
||||
ggufSelectionOverrideMatchesIntent(
|
||||
preferredFile,
|
||||
preferredFileIntent,
|
||||
selectedVariantState.preferredFile,
|
||||
selectedVariantState.preferredFileIntent,
|
||||
)
|
||||
? selectedVariantState.quant
|
||||
: preferredQuant;
|
||||
const selectedQuant =
|
||||
selectedVariantOverride &&
|
||||
sortedVariants?.some((variant) =>
|
||||
|
|
@ -506,6 +532,9 @@ export function LocalOnDeviceCard({
|
|||
setSelectedVariantState({
|
||||
key: variantKey,
|
||||
quant: variant.quant,
|
||||
|
||||
preferredFile,
|
||||
preferredFileIntent,
|
||||
});
|
||||
setVariantOpen(false);
|
||||
}}
|
||||
|
|
|
|||
|
|
@ -411,6 +411,9 @@ export const ModelInspector = memo(function ModelInspector({
|
|||
model,
|
||||
runtime,
|
||||
actions,
|
||||
preferredGgufFile = null,
|
||||
|
||||
preferredGgufFileIntent = 0,
|
||||
isDataset = false,
|
||||
metadataUnavailable = false,
|
||||
selectionHiddenByFilters = false,
|
||||
|
|
@ -419,6 +422,9 @@ export const ModelInspector = memo(function ModelInspector({
|
|||
isDataset?: boolean;
|
||||
metadataUnavailable?: boolean;
|
||||
selectionHiddenByFilters?: boolean;
|
||||
preferredGgufFile?: string | null;
|
||||
|
||||
preferredGgufFileIntent?: number;
|
||||
runtime: ModelInspectorRuntime;
|
||||
actions: ModelInspectorActions;
|
||||
}) {
|
||||
|
|
@ -696,6 +702,9 @@ export const ModelInspector = memo(function ModelInspector({
|
|||
loadingPhase={loadingPhase}
|
||||
gpuGb={gpuGb}
|
||||
systemRamGb={systemRamGb}
|
||||
|
||||
preferredFile={preferredGgufFile}
|
||||
preferredFileIntent={preferredGgufFileIntent}
|
||||
unsupportedReason={
|
||||
unslothSupport.status === "unsupported"
|
||||
? (unslothSupport.reason ?? "Unsupported format")
|
||||
|
|
@ -721,6 +730,9 @@ export const ModelInspector = memo(function ModelInspector({
|
|||
canRun={canRunModel}
|
||||
isActive={isActive}
|
||||
activeQuant={isActive ? (activeGgufVariant ?? null) : null}
|
||||
preferredGgufFile={preferredGgufFile}
|
||||
|
||||
preferredGgufFileIntent={preferredGgufFileIntent}
|
||||
isLoadingThisModel={isLoadingThisModel}
|
||||
gpuGb={gpuGb}
|
||||
systemRamGb={systemRamGb}
|
||||
|
|
|
|||
|
|
@ -365,7 +365,9 @@ export function ModelsPage() {
|
|||
const deviceType = usePlatformStore((s) => s.deviceType);
|
||||
const hubSearch = useSearch({ from: "/hub" });
|
||||
const urlModel = hubSearch.model ?? null;
|
||||
const preferredGgufFile = hubSearch.file ?? null;
|
||||
|
||||
const preferredGgufFileIntent = hubSearch.intent ?? 0;
|
||||
const { selectModel, loadingModel, loadProgress, ejectModel } =
|
||||
useChatModelRuntime();
|
||||
const checkpoint = useChatRuntimeStore((s) => s.params.checkpoint);
|
||||
|
|
@ -1081,7 +1083,7 @@ export function ModelsPage() {
|
|||
setSelected(id);
|
||||
void navigate({
|
||||
to: "/hub",
|
||||
search: (prev) => ({ ...prev, model: id }),
|
||||
search: (prev) => ({ ...prev, model: id, file: undefined }),
|
||||
});
|
||||
},
|
||||
[setSelected, navigate],
|
||||
|
|
@ -1167,7 +1169,7 @@ export function ModelsPage() {
|
|||
setSelected(firstId);
|
||||
void navigate({
|
||||
to: "/hub",
|
||||
search: (prev) => ({ ...prev, model: firstId }),
|
||||
search: (prev) => ({ ...prev, model: firstId, file: undefined }),
|
||||
replace: true,
|
||||
});
|
||||
}, [
|
||||
|
|
@ -1868,6 +1870,9 @@ export function ModelsPage() {
|
|||
>
|
||||
<HubDetailView
|
||||
model={selectedModel}
|
||||
preferredGgufFile={preferredGgufFile}
|
||||
|
||||
preferredGgufFileIntent={preferredGgufFileIntent}
|
||||
isDataset={isDatasetMode}
|
||||
metadataUnavailable={metadataUnavailable}
|
||||
selectionHiddenByFilters={selectionHiddenByFilters}
|
||||
|
|
@ -1890,6 +1895,9 @@ export function ModelsPage() {
|
|||
>
|
||||
<HubDetailView
|
||||
model={selectedModel}
|
||||
preferredGgufFile={preferredGgufFile}
|
||||
|
||||
preferredGgufFileIntent={preferredGgufFileIntent}
|
||||
isDataset={isDatasetMode}
|
||||
metadataUnavailable={metadataUnavailable}
|
||||
selectionHiddenByFilters={selectionHiddenByFilters}
|
||||
|
|
|
|||
33
studio/frontend/src/features/hub/lib/gguf-filename.ts
Normal file
33
studio/frontend/src/features/hub/lib/gguf-filename.ts
Normal file
|
|
@ -0,0 +1,33 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
const GGUF_SPLIT_SUFFIX = /-\d{3,}-of-\d{3,}(?=\.gguf$)/i;
|
||||
|
||||
function normalizeGgufFilename(filename: string): string {
|
||||
return filename
|
||||
.trim()
|
||||
.replace(/\\/g, "/")
|
||||
.replace(GGUF_SPLIT_SUFFIX, "")
|
||||
.toLowerCase();
|
||||
}
|
||||
|
||||
export function ggufFilenamesMatch(
|
||||
left: string | null | undefined,
|
||||
right: string | null | undefined,
|
||||
): boolean {
|
||||
if (!(left && right)) return false;
|
||||
return normalizeGgufFilename(left) === normalizeGgufFilename(right);
|
||||
}
|
||||
|
||||
export function ggufSelectionOverrideMatchesIntent(
|
||||
preferredFile: string | null | undefined,
|
||||
preferredFileIntent: number,
|
||||
selectedPreferredFile: string | null | undefined,
|
||||
selectedPreferredFileIntent: number | undefined,
|
||||
): boolean {
|
||||
return (
|
||||
!preferredFile ||
|
||||
(selectedPreferredFile === preferredFile &&
|
||||
selectedPreferredFileIntent === preferredFileIntent)
|
||||
);
|
||||
}
|
||||
|
|
@ -47,6 +47,8 @@ import {
|
|||
MAX_SEQ_LENGTH_MIN,
|
||||
MAX_SEQ_LENGTH_STEP,
|
||||
MTP_SPECULATIVE_TYPES,
|
||||
N_PARALLEL_MAX,
|
||||
N_PARALLEL_MIN,
|
||||
type PerModelConfig,
|
||||
SPECULATIVE_TYPES,
|
||||
deletePerModelConfig,
|
||||
|
|
@ -91,6 +93,7 @@ function hasNonDefaultAdvanced(config: PerModelConfig): boolean {
|
|||
config.kvCacheDtype != null ||
|
||||
(config.speculativeType ?? "auto") !== "auto" ||
|
||||
config.specDraftNMax != null ||
|
||||
config.nParallel != null ||
|
||||
config.tensorParallel ||
|
||||
config.chatTemplateOverride != null ||
|
||||
(config.gpuMemoryMode ?? "auto") !== "auto" ||
|
||||
|
|
@ -545,6 +548,44 @@ function GgufAdvancedSettings({
|
|||
</div>
|
||||
)}
|
||||
|
||||
<div className={ROW_CLASS}>
|
||||
<div className="flex min-w-0 items-center gap-1.5">
|
||||
<span className={LABEL_CLASS}>Parallel Slots</span>
|
||||
<InfoHint>
|
||||
llama-server decode slots (--parallel) for concurrent requests.
|
||||
Leave blank for the server default. More slots share the context
|
||||
pool and use more VRAM; if they don't fit on GPU, fewer slots are
|
||||
launched.
|
||||
</InfoHint>
|
||||
</div>
|
||||
<input
|
||||
type="number"
|
||||
min={N_PARALLEL_MIN}
|
||||
max={N_PARALLEL_MAX}
|
||||
step={1}
|
||||
value={config.nParallel ?? ""}
|
||||
placeholder="auto"
|
||||
onChange={(event) => {
|
||||
const raw = event.target.value;
|
||||
if (raw === "") {
|
||||
update({ nParallel: null });
|
||||
return;
|
||||
}
|
||||
const parsed = Number.parseInt(raw, 10);
|
||||
if (Number.isFinite(parsed)) {
|
||||
update({
|
||||
nParallel: Math.max(
|
||||
N_PARALLEL_MIN,
|
||||
Math.min(N_PARALLEL_MAX, parsed),
|
||||
),
|
||||
});
|
||||
}
|
||||
}}
|
||||
aria-label="Parallel decode slots"
|
||||
className={NUMBER_INPUT_CLASS}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div className={ROW_CLASS}>
|
||||
<div className="flex min-w-0 items-center gap-1.5">
|
||||
<span className={LABEL_CLASS}>Tensor Parallelism</span>
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ export function useActiveModelConfig(): ActiveModelConfigState {
|
|||
const kvCacheDtype = useChatRuntimeStore((s) => s.kvCacheDtype);
|
||||
const speculativeType = useChatRuntimeStore((s) => s.speculativeType);
|
||||
const specDraftNMax = useChatRuntimeStore((s) => s.specDraftNMax);
|
||||
const nParallel = useChatRuntimeStore((s) => s.nParallel);
|
||||
const tensorParallel = useChatRuntimeStore((s) => s.tensorParallel);
|
||||
const chatTemplateOverride = useChatRuntimeStore(
|
||||
(s) => s.chatTemplateOverride,
|
||||
|
|
@ -44,6 +45,7 @@ export function useActiveModelConfig(): ActiveModelConfigState {
|
|||
kvCacheDtype: kvCacheDtype ?? null,
|
||||
speculativeType: speculativeType ?? "auto",
|
||||
specDraftNMax: specDraftNMax ?? null,
|
||||
nParallel: nParallel ?? null,
|
||||
tensorParallel: tensorParallel ?? false,
|
||||
chatTemplateOverride: chatTemplateOverride ?? null,
|
||||
};
|
||||
|
|
@ -65,6 +67,7 @@ export function useActiveModelConfig(): ActiveModelConfigState {
|
|||
kvCacheDtype,
|
||||
speculativeType,
|
||||
specDraftNMax,
|
||||
nParallel,
|
||||
tensorParallel,
|
||||
chatTemplateOverride,
|
||||
gpuMemoryMode,
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ export function applyPerModelConfigToRuntime(config: PerModelConfig): void {
|
|||
normalizeSpeculativeType(config.speculativeType) ??
|
||||
readPersistedSpeculativeType(),
|
||||
specDraftNMax: config.specDraftNMax ?? null,
|
||||
nParallel: config.nParallel ?? null,
|
||||
tensorParallel: config.tensorParallel ?? false,
|
||||
chatTemplateOverride: cleanTemplate(config.chatTemplateOverride),
|
||||
// GPU Memory knobs are per-model (GGUF-only). Absent = defaults; the mode is
|
||||
|
|
@ -82,6 +83,7 @@ export function currentRuntimePerModelConfig(
|
|||
kvCacheDtype: s.kvCacheDtype ?? null,
|
||||
speculativeType: normalizeSpeculativeType(s.speculativeType),
|
||||
specDraftNMax: s.specDraftNMax ?? null,
|
||||
nParallel: s.nParallel ?? null,
|
||||
tensorParallel: s.tensorParallel ?? false,
|
||||
chatTemplateOverride: cleanTemplate(s.chatTemplateOverride),
|
||||
// Snapshot the live GPU knobs too so a failed switch rolls the previous
|
||||
|
|
@ -106,6 +108,7 @@ export function perModelConfigsEqual(
|
|||
normalizeSpeculativeType(a.speculativeType) ===
|
||||
normalizeSpeculativeType(b.speculativeType) &&
|
||||
(a.specDraftNMax ?? null) === (b.specDraftNMax ?? null) &&
|
||||
(a.nParallel ?? null) === (b.nParallel ?? null) &&
|
||||
Boolean(a.tensorParallel) === Boolean(b.tensorParallel) &&
|
||||
cleanTemplate(a.chatTemplateOverride) ===
|
||||
cleanTemplate(b.chatTemplateOverride) &&
|
||||
|
|
|
|||
|
|
@ -54,6 +54,7 @@ export function loadedConfigSignature(
|
|||
config.kvCacheDtype ?? "",
|
||||
config.speculativeType ?? "",
|
||||
config.specDraftNMax ?? "",
|
||||
config.nParallel ?? "",
|
||||
config.tensorParallel ? "1" : "0",
|
||||
config.chatTemplateOverride == null
|
||||
? ""
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ export interface PerModelConfig {
|
|||
kvCacheDtype: string | null;
|
||||
speculativeType: string | null;
|
||||
specDraftNMax: number | null;
|
||||
nParallel: number | null;
|
||||
tensorParallel: boolean;
|
||||
chatTemplateOverride: string | null;
|
||||
// GPU Memory controls (per-model, GGUF-only), optional so older blobs still
|
||||
|
|
@ -33,10 +34,16 @@ export const DEFAULT_PER_MODEL_CONFIG: PerModelConfig = {
|
|||
kvCacheDtype: null,
|
||||
speculativeType: null,
|
||||
specDraftNMax: null,
|
||||
nParallel: null,
|
||||
tensorParallel: false,
|
||||
chatTemplateOverride: null,
|
||||
};
|
||||
|
||||
// Mirrors llama_server_args.py PARALLEL_MIN/MAX (LoadRequest.n_parallel
|
||||
// bounds). null = follow the server-wide default.
|
||||
export const N_PARALLEL_MIN = 1;
|
||||
export const N_PARALLEL_MAX = 64;
|
||||
|
||||
export const MAX_SEQ_LENGTH_MIN = 128;
|
||||
export const MAX_SEQ_LENGTH_MAX = 1048576;
|
||||
export const MAX_SEQ_LENGTH_STEP = 128;
|
||||
|
|
@ -92,6 +99,7 @@ const STORED_CONFIG_FIELDS = new Set([
|
|||
"kvCacheDtype",
|
||||
"speculativeType",
|
||||
"specDraftNMax",
|
||||
"nParallel",
|
||||
"tensorParallel",
|
||||
"chatTemplateOverride",
|
||||
"gpuMemoryMode",
|
||||
|
|
@ -295,6 +303,8 @@ function legacyEntryToConfig(raw: Record<string, unknown>): PerModelConfig {
|
|||
typeof raw.speculativeType === "string" ? raw.speculativeType : null,
|
||||
specDraftNMax:
|
||||
typeof raw.specDraftNMax === "number" ? raw.specDraftNMax : null,
|
||||
// Legacy blobs predate the parallel-slots knob.
|
||||
nParallel: null,
|
||||
tensorParallel:
|
||||
typeof raw.tensorParallel === "boolean" ? raw.tensorParallel : false,
|
||||
chatTemplateOverride: null,
|
||||
|
|
@ -466,6 +476,10 @@ function normalizeV1(partial: RawConfig): PerModelConfig {
|
|||
: null,
|
||||
speculativeType,
|
||||
specDraftNMax,
|
||||
nParallel:
|
||||
typeof partial.nParallel === "number" && Number.isFinite(partial.nParallel)
|
||||
? Math.max(N_PARALLEL_MIN, Math.min(N_PARALLEL_MAX, Math.round(partial.nParallel)))
|
||||
: null,
|
||||
tensorParallel:
|
||||
typeof partial.tensorParallel === "boolean"
|
||||
? partial.tensorParallel
|
||||
|
|
@ -613,6 +627,7 @@ export function isDefaultConfig(config: PerModelConfig): boolean {
|
|||
(config.kvCacheDtype ?? null) === DEFAULT_PER_MODEL_CONFIG.kvCacheDtype &&
|
||||
config.speculativeType === DEFAULT_PER_MODEL_CONFIG.speculativeType &&
|
||||
config.specDraftNMax == null &&
|
||||
config.nParallel == null &&
|
||||
Boolean(config.tensorParallel) ===
|
||||
Boolean(DEFAULT_PER_MODEL_CONFIG.tensorParallel) &&
|
||||
(config.chatTemplateOverride ?? null) === null &&
|
||||
|
|
|
|||
130
studio/frontend/tests/artifact-source-key.test.ts
Normal file
130
studio/frontend/tests/artifact-source-key.test.ts
Normal file
|
|
@ -0,0 +1,130 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
import assert from "node:assert/strict";
|
||||
import { readFileSync } from "node:fs";
|
||||
import test from "node:test";
|
||||
import { fileURLToPath } from "node:url";
|
||||
|
||||
import ts from "typescript";
|
||||
|
||||
import {
|
||||
buildArtifactSourceKey,
|
||||
createArtifactId,
|
||||
createChatArtifact,
|
||||
hashArtifactCode,
|
||||
} from "../src/features/chat/artifacts/types.ts";
|
||||
|
||||
// The shipped helper the component keys on, not a copy of it.
|
||||
const sourceKey = buildArtifactSourceKey;
|
||||
|
||||
const toolInput = (code: string) => ({
|
||||
code,
|
||||
source: "tool" as const,
|
||||
threadId: "thread-1",
|
||||
sourceMessageId: "msg-1",
|
||||
sourceToolCallId: "call_0",
|
||||
});
|
||||
|
||||
const fenceInput = (code: string) => ({
|
||||
code,
|
||||
source: "fence" as const,
|
||||
threadId: "thread-1",
|
||||
sourceMessageId: "msg-1",
|
||||
});
|
||||
|
||||
test("tool artifact IDs are stable across code changes, so the ID alone is not enough", () => {
|
||||
const first = createArtifactId(toolInput("<p>first</p>"));
|
||||
const second = createArtifactId(toolInput("<p>second</p>"));
|
||||
assert.equal(first, second);
|
||||
});
|
||||
|
||||
test("the source key changes when a tool artifact's code changes", () => {
|
||||
const first = createChatArtifact(toolInput("<p>first</p>"));
|
||||
const second = createChatArtifact(toolInput("<p>second</p>"));
|
||||
assert.notEqual(sourceKey(first), sourceKey(second));
|
||||
});
|
||||
|
||||
test("the source key changes when switching between fence artifacts", () => {
|
||||
const first = createChatArtifact(fenceInput("<p>alpha</p>"));
|
||||
const second = createChatArtifact(fenceInput("<p>bravo</p>"));
|
||||
assert.notEqual(sourceKey(first), sourceKey(second));
|
||||
});
|
||||
|
||||
test("the source key is stable for an unchanged artifact, so no needless remount", () => {
|
||||
const code = "<p>same</p>";
|
||||
assert.equal(
|
||||
sourceKey(createChatArtifact(toolInput(code))),
|
||||
sourceKey(createChatArtifact(toolInput(code))),
|
||||
);
|
||||
});
|
||||
|
||||
// Equal line count, the shape where Streamdown's comparator sees no change.
|
||||
test("the source key changes for two canvases with the same shape", () => {
|
||||
const first = createChatArtifact(
|
||||
toolInput("<html>\n<body>\n<h1>Alpha</h1>\n</body>\n</html>"),
|
||||
);
|
||||
const second = createChatArtifact(
|
||||
toolInput("<html>\n<body>\n<h1>Bravo</h1>\n</body>\n</html>"),
|
||||
);
|
||||
assert.equal(first.code.length, second.code.length);
|
||||
assert.equal(first.code.split("\n").length, second.code.split("\n").length);
|
||||
assert.notEqual(sourceKey(first), sourceKey(second));
|
||||
});
|
||||
|
||||
test("hashArtifactCode separates same-length codes and empty from whitespace", () => {
|
||||
assert.notEqual(hashArtifactCode("<p>ab</p>"), hashArtifactCode("<p>ba</p>"));
|
||||
assert.notEqual(hashArtifactCode(""), hashArtifactCode(" "));
|
||||
});
|
||||
|
||||
const KEYED_BY_HELPER = /^\{buildArtifactSourceKey\(\s*artifact\s*\)\}$/;
|
||||
|
||||
const SURFACE_PATH = fileURLToPath(
|
||||
new URL(
|
||||
"../src/features/chat/artifacts/artifact-surface.tsx",
|
||||
import.meta.url,
|
||||
),
|
||||
);
|
||||
|
||||
/** The opening tag of `node`, for both `<x>` and `<x />`. */
|
||||
const openingTag = (node: ts.Node): ts.JsxOpeningLikeElement | null => {
|
||||
if (ts.isJsxSelfClosingElement(node)) return node;
|
||||
if (ts.isJsxElement(node)) return node.openingElement;
|
||||
return null;
|
||||
};
|
||||
|
||||
/** The `key` expression on the source view's Streamdown, or null if unkeyed. */
|
||||
function readStreamdownKey(): string | null {
|
||||
const source = ts.createSourceFile(
|
||||
SURFACE_PATH,
|
||||
readFileSync(SURFACE_PATH, "utf8"),
|
||||
ts.ScriptTarget.ESNext,
|
||||
true,
|
||||
ts.ScriptKind.TSX,
|
||||
);
|
||||
let key: string | null = null;
|
||||
const visit = (node: ts.Node): void => {
|
||||
const opening = openingTag(node);
|
||||
if (opening?.tagName.getText() === "Streamdown") {
|
||||
for (const attribute of opening.attributes.properties) {
|
||||
if (
|
||||
ts.isJsxAttribute(attribute) &&
|
||||
attribute.name.getText() === "key"
|
||||
) {
|
||||
key = attribute.initializer?.getText() ?? "";
|
||||
}
|
||||
}
|
||||
}
|
||||
node.forEachChild(visit);
|
||||
};
|
||||
source.forEachChild(visit);
|
||||
return key;
|
||||
}
|
||||
|
||||
// Without this the suite passes with the key deleted, which is the regression.
|
||||
// No DOM renderer is available here, so assert the wiring in the source.
|
||||
test("the source view's Streamdown is keyed by the shipped helper", () => {
|
||||
const key = readStreamdownKey();
|
||||
assert.ok(key, "source view <Streamdown> has no key prop");
|
||||
assert.match(key, KEYED_BY_HELPER);
|
||||
});
|
||||
|
|
@ -17,6 +17,7 @@ const LIVE: PerModelConfig = {
|
|||
kvCacheDtype: "q8_0",
|
||||
speculativeType: "ngram",
|
||||
specDraftNMax: 6,
|
||||
nParallel: 4,
|
||||
tensorParallel: true,
|
||||
chatTemplateOverride: null,
|
||||
gpuMemoryMode: "manual",
|
||||
|
|
@ -32,6 +33,7 @@ const SAVED: PerModelConfig = {
|
|||
kvCacheDtype: null,
|
||||
speculativeType: "auto",
|
||||
specDraftNMax: null,
|
||||
nParallel: null,
|
||||
tensorParallel: false,
|
||||
chatTemplateOverride: null,
|
||||
gpuMemoryMode: "auto",
|
||||
|
|
@ -100,6 +102,7 @@ test("every mirrored setting moves the instance key", () => {
|
|||
{ ...LIVE, kvCacheDtype: "f16" },
|
||||
{ ...LIVE, speculativeType: "off" },
|
||||
{ ...LIVE, specDraftNMax: 4 },
|
||||
{ ...LIVE, nParallel: 1 },
|
||||
{ ...LIVE, tensorParallel: false },
|
||||
{ ...LIVE, chatTemplateOverride: "{{ bos_token }}" },
|
||||
{ ...LIVE, gpuMemoryMode: "auto" },
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ function config(maxSeqLength: number, kvCacheDtype: string | null = null) {
|
|||
kvCacheDtype,
|
||||
speculativeType: null,
|
||||
specDraftNMax: null,
|
||||
nParallel: null,
|
||||
tensorParallel: false,
|
||||
chatTemplateOverride: null,
|
||||
};
|
||||
|
|
|
|||
45
studio/frontend/tests/tool-status.test.ts
Normal file
45
studio/frontend/tests/tool-status.test.ts
Normal file
|
|
@ -0,0 +1,45 @@
|
|||
// SPDX-License-Identifier: AGPL-3.0-only
|
||||
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
import assert from "node:assert/strict";
|
||||
import test from "node:test";
|
||||
|
||||
import {
|
||||
NUDGE_TOOL_CALLS_STATUS,
|
||||
toolStatusKind,
|
||||
} from "../src/features/chat/utils/tool-status.ts";
|
||||
|
||||
test("the nudge status is the exact string the backend sends", () => {
|
||||
// Mirrors tool_call_parser.py, so a reword on either side must break here.
|
||||
assert.equal(NUDGE_TOOL_CALLS_STATUS, "Nudging tool calls");
|
||||
assert.equal(toolStatusKind(NUDGE_TOOL_CALLS_STATUS), "nudge");
|
||||
});
|
||||
|
||||
test("sandbox tools keep the terminal glyph", () => {
|
||||
for (const status of [
|
||||
"Running Python: print(1)",
|
||||
"Running Python...",
|
||||
"Running: ls -la",
|
||||
"Running command...",
|
||||
]) {
|
||||
assert.equal(toolStatusKind(status), "terminal", status);
|
||||
}
|
||||
});
|
||||
|
||||
test("every other status keeps the globe", () => {
|
||||
for (const status of [
|
||||
"Searching: red square",
|
||||
"Reading: unsloth.ai",
|
||||
"Reading page...",
|
||||
"Searching documents: quarterly report",
|
||||
"Calling: get_weather",
|
||||
]) {
|
||||
assert.equal(toolStatusKind(status), "web", status);
|
||||
}
|
||||
});
|
||||
|
||||
test("a status that merely mentions nudging is not the nudge itself", () => {
|
||||
// Exact match only: a tool named after the phrase must not steal the spinner.
|
||||
assert.equal(toolStatusKind("Calling: Nudging tool calls"), "web");
|
||||
assert.equal(toolStatusKind("Nudging tool calls again"), "web");
|
||||
});
|
||||
|
|
@ -4807,6 +4807,15 @@ def _native_linux_system_rocm_lib_dirs(binary_dir: str = "") -> list[str]:
|
|||
os.path.join(d, "libhsa-runtime64.so.1")
|
||||
):
|
||||
out.append(d)
|
||||
# ROCm keeps LLVM's versioned runtime under <root>/lib/llvm, so a
|
||||
# lib64 host still finds it under lib. Probe both and keep them
|
||||
# ahead of the bundle, else system libamd_comgr binds to the
|
||||
# bundle's incompatible libLLVM.so.*.
|
||||
for _sub in (lib_sub, "lib"):
|
||||
llvm_lib = os.path.join(base, _sub, "llvm", "lib")
|
||||
if llvm_lib not in seen and os.path.isdir(llvm_lib):
|
||||
seen.add(llvm_lib)
|
||||
out.append(llvm_lib)
|
||||
return out
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -426,8 +426,8 @@ _GFX_TO_AMD_INDEX_ARCH: dict[str, str] = {
|
|||
}
|
||||
|
||||
# bitsandbytes continuous-release_main wheels with the ROCm 4-bit GEMV fix
|
||||
# (bnb PR #1887, post-0.49.2). bnb <= 0.49.2 NaNs at decode shape on every
|
||||
# AMD GPU. Drop the pin once bnb 0.50+ ships on PyPI.
|
||||
# (bnb #1887, post-0.49.2). bnb <= 0.49.2 NaNs at decode shape on every AMD GPU;
|
||||
# PyPI 0.50.0 is the first release with the fix, so the fallback below is safe.
|
||||
_BNB_ROCM_PRERELEASE_URLS: dict[str, str] = {
|
||||
"x86_64": (
|
||||
"https://github.com/bitsandbytes-foundation/bitsandbytes/releases/"
|
||||
|
|
@ -448,7 +448,8 @@ _BNB_ROCM_PRERELEASE_URLS: dict[str, str] = {
|
|||
"bitsandbytes-1.33.7.preview-py3-none-win_amd64.whl"
|
||||
),
|
||||
}
|
||||
_BNB_ROCM_PYPI_FALLBACK = "bitsandbytes>=0.49.1"
|
||||
# Keep in step with the amd extra in pyproject.toml and the install.sh fallback.
|
||||
_BNB_ROCM_PYPI_FALLBACK = "bitsandbytes>=0.50.0"
|
||||
|
||||
|
||||
def _bnb_rocm_prerelease_url() -> str | None:
|
||||
|
|
@ -460,6 +461,16 @@ def _bnb_rocm_prerelease_url() -> str | None:
|
|||
return _BNB_ROCM_PRERELEASE_URLS.get(arch)
|
||||
|
||||
|
||||
def _bnb_rocm_arch_has_binary() -> bool:
|
||||
"""False on aarch64: bitsandbytes ships no ROCm kernels there at any version.
|
||||
The PyPI 0.50.0 and continuous-release_main aarch64 wheels both carry only
|
||||
libbitsandbytes_cpu.so plus CUDA variants, so neither install path gives
|
||||
aarch64 a 4-bit backend and neither message may claim one.
|
||||
"""
|
||||
arch = platform.machine().lower()
|
||||
return {"amd64": "x86_64", "arm64": "aarch64"}.get(arch, arch) != "aarch64"
|
||||
|
||||
|
||||
def _amd_smi_env() -> dict[str, str] | None:
|
||||
"""On Windows, env with __COMPAT_LAYER=RunAsInvoker; None elsewhere.
|
||||
NB: RunAsInvoker doesn't stop amd-smi's runtime elevation (its manifest is
|
||||
|
|
@ -1243,29 +1254,46 @@ _rocm_windows_torch_installed: bool = False
|
|||
|
||||
|
||||
def _install_bnb_windows_rocm() -> bool:
|
||||
"""Install the AMD Windows BNB prerelease wheel. Returns True on success.
|
||||
"""Install AMD Windows BNB, pre-release wheel first. Returns True on success.
|
||||
|
||||
The continuous-release wheel is intentionally mismatched: the filename
|
||||
encodes 1.33.7.preview (parsed as 1.33.7rc0 by PEP 440) while the wheel
|
||||
metadata reports 0.50.0.dev0. uv rejects this filename/metadata mismatch,
|
||||
and bypassing it with UV_SKIP_WHEEL_FILENAME_CHECK still leaves uv mangling
|
||||
the bitsandbytes install. Per the AMD install guide
|
||||
(https://unsloth.ai/docs/get-started/install/amd/amd-hackathon) the wheel
|
||||
must be installed with plain pip, not uv, so we force pip (force_pip=True);
|
||||
plain pip performs no wheel filename/metadata check.
|
||||
The wheel's filename version (1.33.7.preview, PEP 440 1.33.7rc0) does not
|
||||
match its metadata (0.50.x.dev0). uv rejects the mismatch and still mangles
|
||||
the install under UV_SKIP_WHEEL_FILENAME_CHECK, so force plain pip, which
|
||||
performs no such check. Per the AMD install guide
|
||||
(https://unsloth.ai/docs/get-started/install/amd/amd-hackathon).
|
||||
|
||||
When that URL is blocked, fall back to PyPI. Its win_amd64 wheel ships
|
||||
libbitsandbytes_rocm{714,72}.dll from 0.50.0 on, so the fallback is a real
|
||||
ROCm build; before 0.50.0 it was CUDA-only, which is why there was none.
|
||||
"""
|
||||
_bnb_win_url = _BNB_ROCM_PRERELEASE_URLS.get("win_amd64")
|
||||
if _bnb_win_url is None:
|
||||
return False
|
||||
_ok = pip_install_try(
|
||||
"bitsandbytes (AMD Windows, pre-release main)",
|
||||
"--force-reinstall",
|
||||
"--no-cache-dir",
|
||||
"--no-deps",
|
||||
_bnb_win_url,
|
||||
constrain = False,
|
||||
force_pip = True,
|
||||
)
|
||||
_ok = False
|
||||
if _bnb_win_url is not None:
|
||||
_ok = pip_install_try(
|
||||
"bitsandbytes (AMD Windows, pre-release main)",
|
||||
"--force-reinstall",
|
||||
"--no-cache-dir",
|
||||
"--no-deps",
|
||||
_bnb_win_url,
|
||||
constrain = False,
|
||||
force_pip = True,
|
||||
)
|
||||
if not _ok:
|
||||
print(
|
||||
_red(
|
||||
" bnb pre-release install failed; falling back to PyPI "
|
||||
f"{_BNB_ROCM_PYPI_FALLBACK}, which carries the ROCm 4-bit fix"
|
||||
)
|
||||
)
|
||||
if not _ok:
|
||||
_ok = pip_install_try(
|
||||
"bitsandbytes (AMD Windows)",
|
||||
"--force-reinstall",
|
||||
"--no-cache-dir",
|
||||
"--no-deps",
|
||||
_BNB_ROCM_PYPI_FALLBACK,
|
||||
constrain = False,
|
||||
)
|
||||
if not _ok:
|
||||
return False
|
||||
# Detect the actual ROCm DLL suffix in the wheel and set BNB_ROCM_VERSION so bnb
|
||||
|
|
@ -1755,8 +1783,8 @@ def _ensure_rocm_torch() -> None:
|
|||
pass
|
||||
if _torch_ok:
|
||||
_rocm_windows_torch_installed = True
|
||||
# ROCm torch is already installed, but the AMD Windows BNB wheel is still
|
||||
# needed (the PyPI bitsandbytes ships only CUDA DLLs, fails on ROCm).
|
||||
# ROCm torch is already installed, but bnb still needs the ROCm build
|
||||
# (pre-release wheel, else PyPI >=0.50.0).
|
||||
_install_bnb_windows_rocm()
|
||||
return
|
||||
# torch was wiped between runs; fall through to the full install path
|
||||
|
|
@ -1834,12 +1862,12 @@ def _ensure_rocm_torch() -> None:
|
|||
# separate dependency -- a BNB install failure must NOT roll back the
|
||||
# torch ROCm install.
|
||||
_rocm_windows_torch_installed = True
|
||||
# Always install AMD Windows bitsandbytes -- the PyPI wheel ships only
|
||||
# CUDA DLLs and fails on ROCm. Install even when torch was already a
|
||||
# ROCm build so `studio update` repairs a broken bnb.
|
||||
# Always install AMD Windows bitsandbytes, even when torch was already a
|
||||
# ROCm build, so `studio update` repairs a broken bnb.
|
||||
if not _install_bnb_windows_rocm():
|
||||
print(
|
||||
" Warning: AMD Windows bitsandbytes install failed; "
|
||||
" Warning: AMD Windows bitsandbytes install failed "
|
||||
"(pre-release and PyPI); "
|
||||
"ROCm torch is installed but bitsandbytes may need manual install"
|
||||
)
|
||||
return
|
||||
|
|
@ -2170,10 +2198,13 @@ def _ensure_rocm_torch() -> None:
|
|||
force_pip = True,
|
||||
)
|
||||
if not _bnb_installed:
|
||||
_fallback_note = (
|
||||
", which carries the ROCm 4-bit fix" if _bnb_rocm_arch_has_binary() else ""
|
||||
)
|
||||
print(
|
||||
_red(
|
||||
" bnb pre-release install failed; falling back to PyPI "
|
||||
"(4-bit decode will be broken on ROCm)"
|
||||
f"{_BNB_ROCM_PYPI_FALLBACK}{_fallback_note}"
|
||||
)
|
||||
)
|
||||
if not _bnb_installed:
|
||||
|
|
@ -2185,6 +2216,14 @@ def _ensure_rocm_torch() -> None:
|
|||
_BNB_ROCM_PYPI_FALLBACK,
|
||||
constrain = False,
|
||||
)
|
||||
if not _bnb_rocm_arch_has_binary():
|
||||
print(
|
||||
_red(
|
||||
" aarch64: bitsandbytes ships no ROCm kernels on this arch; "
|
||||
"4-bit QLoRA needs a source build -- "
|
||||
"https://docs.unsloth.ai/get-started/install-and-update/amd"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# _uv_safe_path is imported from backend.utils.uv_path_safety (shared with mlx_repair).
|
||||
|
|
|
|||
|
|
@ -3122,7 +3122,7 @@ sys.exit(0 if install_manifest.remove_manifest() else 1)
|
|||
if (-not $_ManifestDropped) {
|
||||
Write-Host "[ERROR] Could not remove the stale unsloth_install_manifest.json." -ForegroundColor Red
|
||||
Write-Host " Refusing to install behind a marker that still reports this venv as complete." -ForegroundColor Red
|
||||
exit 1
|
||||
Exit-SetupFailure "Could not remove the stale unsloth_install_manifest.json"
|
||||
}
|
||||
|
||||
if ($script:UnslothVerbose) {
|
||||
|
|
|
|||
126
studio/src-tauri/Cargo.lock
generated
126
studio/src-tauri/Cargo.lock
generated
|
|
@ -558,6 +558,26 @@ version = "0.10.2"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a6ef517f0926dd24a1582492c791b6a4818a4d94e789a334894aa15b0d12f55c"
|
||||
|
||||
[[package]]
|
||||
name = "const-random"
|
||||
version = "0.1.18"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "87e00182fe74b066627d63b85fd550ac2998d4b0bd86bfed477a0ae4c7c71359"
|
||||
dependencies = [
|
||||
"const-random-macro",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "const-random-macro"
|
||||
version = "0.1.16"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f9d839f2a20b0aee515dc581a6172f2321f96cab76c1a38a4c584a194955390e"
|
||||
dependencies = [
|
||||
"getrandom 0.2.17",
|
||||
"once_cell",
|
||||
"tiny-keccak",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "convert_case"
|
||||
version = "0.4.0"
|
||||
|
|
@ -896,7 +916,7 @@ dependencies = [
|
|||
"libc",
|
||||
"option-ext",
|
||||
"redox_users",
|
||||
"windows-sys 0.59.0",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -945,6 +965,15 @@ dependencies = [
|
|||
"syn 2.0.117",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "dlv-list"
|
||||
version = "0.5.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "442039f5147480ba31067cb00ada1adae6892028e40e45fc5de7b7df6dcc1b5f"
|
||||
dependencies = [
|
||||
"const-random",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "dom_query"
|
||||
version = "0.27.0"
|
||||
|
|
@ -1111,7 +1140,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"windows-sys 0.59.0",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -1738,6 +1767,12 @@ version = "0.12.3"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8a9ee70c43aaf417c914396645a0fa852624801b24ebb7ae78fe8272889ac888"
|
||||
|
||||
[[package]]
|
||||
name = "hashbrown"
|
||||
version = "0.14.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e5274423e17b7c9fc20b6e7e208532f9b19825d82dfd615708b70edd83df41f1"
|
||||
|
||||
[[package]]
|
||||
name = "hashbrown"
|
||||
version = "0.15.5"
|
||||
|
|
@ -1941,7 +1976,7 @@ dependencies = [
|
|||
"tokio",
|
||||
"tower-service",
|
||||
"tracing",
|
||||
"windows-registry",
|
||||
"windows-registry 0.6.1",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -2531,7 +2566,7 @@ dependencies = [
|
|||
"png 0.18.1",
|
||||
"serde",
|
||||
"thiserror 2.0.18",
|
||||
"windows-sys 0.60.2",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -2944,6 +2979,16 @@ version = "0.2.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d"
|
||||
|
||||
[[package]]
|
||||
name = "ordered-multimap"
|
||||
version = "0.7.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "49203cdcae0030493bad186b28da2fa25645fa276a51b6fec8010d281e02ef79"
|
||||
dependencies = [
|
||||
"dlv-list",
|
||||
"hashbrown 0.14.5",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ordered-stream"
|
||||
version = "0.2.0"
|
||||
|
|
@ -2961,7 +3006,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "7d8fae84b431384b68627d0f9b3b1245fcf9f46f6c0e3dc902e9dce64edd1967"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"windows-sys 0.45.0",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -3824,6 +3869,16 @@ dependencies = [
|
|||
"windows-sys 0.52.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rust-ini"
|
||||
version = "0.21.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "796e8d2b6696392a43bea58116b667fb4c29727dc5abd27d6acf338bb4f688c7"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"ordered-multimap",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustc-hash"
|
||||
version = "2.1.1"
|
||||
|
|
@ -3849,7 +3904,7 @@ dependencies = [
|
|||
"errno",
|
||||
"libc",
|
||||
"linux-raw-sys",
|
||||
"windows-sys 0.59.0",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -3905,7 +3960,7 @@ dependencies = [
|
|||
"security-framework",
|
||||
"security-framework-sys",
|
||||
"webpki-root-certs",
|
||||
"windows-sys 0.59.0",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -4347,7 +4402,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "3a766e1110788c36f4fa1c2b71b387a7815aa65f88ce0229841826633d93723e"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"windows-sys 0.60.2",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -4774,6 +4829,27 @@ dependencies = [
|
|||
"thiserror 2.0.18",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tauri-plugin-deep-link"
|
||||
version = "2.4.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "70ee75bc5627f77bfdf40c913255ebc258117b10ebe2b2239a1a1cf40b0b58aa"
|
||||
dependencies = [
|
||||
"dunce",
|
||||
"plist",
|
||||
"rust-ini",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tauri",
|
||||
"tauri-plugin",
|
||||
"tauri-utils",
|
||||
"thiserror 2.0.18",
|
||||
"tracing",
|
||||
"url",
|
||||
"windows-registry 0.5.3",
|
||||
"windows-result 0.3.4",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tauri-plugin-dialog"
|
||||
version = "2.7.1"
|
||||
|
|
@ -4876,6 +4952,7 @@ dependencies = [
|
|||
"serde",
|
||||
"serde_json",
|
||||
"tauri",
|
||||
"tauri-plugin-deep-link",
|
||||
"thiserror 2.0.18",
|
||||
"tracing",
|
||||
"windows-sys 0.60.2",
|
||||
|
|
@ -5014,7 +5091,7 @@ dependencies = [
|
|||
"serde_with",
|
||||
"swift-rs",
|
||||
"thiserror 2.0.18",
|
||||
"toml 0.9.12+spec-1.1.0",
|
||||
"toml 1.1.2+spec-1.1.0",
|
||||
"url",
|
||||
"urlpattern",
|
||||
"uuid",
|
||||
|
|
@ -5051,10 +5128,10 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd"
|
||||
dependencies = [
|
||||
"fastrand",
|
||||
"getrandom 0.3.4",
|
||||
"getrandom 0.4.2",
|
||||
"once_cell",
|
||||
"rustix",
|
||||
"windows-sys 0.59.0",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -5174,6 +5251,15 @@ dependencies = [
|
|||
"time-core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tiny-keccak"
|
||||
version = "2.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2c9d3793400a45f954c52e73d068316d76b6f4e36977e3fcebb13a2721e80237"
|
||||
dependencies = [
|
||||
"crunchy",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tinystr"
|
||||
version = "0.8.2"
|
||||
|
|
@ -5460,7 +5546,7 @@ dependencies = [
|
|||
"png 0.18.1",
|
||||
"serde",
|
||||
"thiserror 2.0.18",
|
||||
"windows-sys 0.60.2",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -5500,7 +5586,7 @@ checksum = "f2f6fb2847f6742cd76af783a2a2c49e9375d0a111c7bef6f71cd9e738c72d6e"
|
|||
dependencies = [
|
||||
"memoffset",
|
||||
"tempfile",
|
||||
"windows-sys 0.60.2",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -5590,6 +5676,7 @@ dependencies = [
|
|||
"tauri",
|
||||
"tauri-build",
|
||||
"tauri-plugin-clipboard-manager",
|
||||
"tauri-plugin-deep-link",
|
||||
"tauri-plugin-dialog",
|
||||
"tauri-plugin-notification",
|
||||
"tauri-plugin-opener",
|
||||
|
|
@ -6069,7 +6156,7 @@ version = "0.1.11"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22"
|
||||
dependencies = [
|
||||
"windows-sys 0.59.0",
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -6257,6 +6344,17 @@ dependencies = [
|
|||
"windows-link 0.2.1",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "windows-registry"
|
||||
version = "0.5.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5b8a9ed28765efc97bbc954883f4e6796c33a06546ebafacbabee9696967499e"
|
||||
dependencies = [
|
||||
"windows-link 0.1.3",
|
||||
"windows-result 0.3.4",
|
||||
"windows-strings 0.4.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "windows-registry"
|
||||
version = "0.6.1"
|
||||
|
|
|
|||
|
|
@ -7,7 +7,8 @@ edition = "2021"
|
|||
|
||||
[dependencies]
|
||||
tauri = { version = "2", features = ["tray-icon"] }
|
||||
tauri-plugin-single-instance = "2"
|
||||
tauri-plugin-single-instance = { version = "2", features = ["deep-link"] }
|
||||
tauri-plugin-deep-link = "2"
|
||||
tauri-plugin-process = "2"
|
||||
serde = { version = "1", features = ["derive"] }
|
||||
serde_json = "1"
|
||||
|
|
|
|||
|
|
@ -16,10 +16,12 @@
|
|||
"core:window:allow-start-dragging",
|
||||
"core:window:allow-start-resize-dragging",
|
||||
"core:window:allow-minimize",
|
||||
"core:window:allow-unminimize",
|
||||
"core:window:allow-toggle-maximize",
|
||||
"core:window:allow-close",
|
||||
"core:tray:default",
|
||||
"process:default",
|
||||
"deep-link:default",
|
||||
"notification:allow-is-permission-granted",
|
||||
"notification:allow-request-permission",
|
||||
"notification:allow-notify",
|
||||
|
|
|
|||
12
studio/src-tauri/linux/unsloth.desktop
Normal file
12
studio/src-tauri/linux/unsloth.desktop
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
[Desktop Entry]
|
||||
Categories={{categories}}
|
||||
{{#if comment}}
|
||||
Comment={{comment}}
|
||||
{{/if}}
|
||||
Exec={{exec}} %u
|
||||
StartupWMClass={{exec}}
|
||||
Icon={{icon}}
|
||||
Name={{name}}
|
||||
Terminal=false
|
||||
Type=Application
|
||||
MimeType=x-scheme-handler/unsloth;
|
||||
|
|
@ -17,7 +17,7 @@ mod process;
|
|||
mod update;
|
||||
mod windows_job;
|
||||
|
||||
use log::info;
|
||||
use log::{info, warn};
|
||||
use process::new_backend_state;
|
||||
use simplelog::{
|
||||
CombinedLogger, Config, LevelFilter, SharedLogger, TermLogger, TerminalMode, WriteLogger,
|
||||
|
|
@ -176,9 +176,11 @@ fn main() {
|
|||
.plugin(tauri_plugin_single_instance::init(|app, _args, _cwd| {
|
||||
if let Some(window) = app.get_webview_window("main") {
|
||||
let _ = window.show();
|
||||
let _ = window.unminimize();
|
||||
let _ = window.set_focus();
|
||||
}
|
||||
}))
|
||||
.plugin(tauri_plugin_deep_link::init())
|
||||
.plugin(tauri_plugin_process::init())
|
||||
.plugin(tauri_plugin_opener::init())
|
||||
.plugin(tauri_plugin_notification::init())
|
||||
|
|
@ -234,6 +236,13 @@ fn main() {
|
|||
has_saved_window_state,
|
||||
])
|
||||
.setup(|app| {
|
||||
#[cfg(any(target_os = "linux", all(debug_assertions, windows)))]
|
||||
{
|
||||
use tauri_plugin_deep_link::DeepLinkExt;
|
||||
if let Err(error) = app.deep_link().register_all() {
|
||||
warn!("Failed to register deep-link handlers: {error}");
|
||||
}
|
||||
}
|
||||
#[cfg(any(target_os = "windows", target_os = "linux"))]
|
||||
setup_custom_titlebar(app)?;
|
||||
setup_tray(app)?;
|
||||
|
|
|
|||
|
|
@ -335,7 +335,13 @@ mod tests {
|
|||
let path = std::env::temp_dir().join(OsString::from_vec(vec![
|
||||
b'u', b'n', b's', b'l', b'o', b't', b'h', 0xff, b'.', b'c', b's', b'v',
|
||||
]));
|
||||
fs::write(&path, "role,content\nuser,hello\n").unwrap();
|
||||
// Linux happily stores arbitrary bytes in a filename, but macOS enforces
|
||||
// UTF-8 on APFS/HFS+ and rejects this name outright. The name-recovery
|
||||
// path being asserted here is only reachable where such a file can
|
||||
// exist, so skip rather than fail on filesystems that forbid it.
|
||||
if fs::write(&path, "role,content\nuser,hello\n").is_err() {
|
||||
return;
|
||||
}
|
||||
let imported = read_selected_import(Some(path.clone())).unwrap().unwrap();
|
||||
assert_eq!(imported.name, "chat-import.csv");
|
||||
let _ = fs::remove_file(path);
|
||||
|
|
|
|||
|
|
@ -704,7 +704,15 @@ mod tests {
|
|||
fs::write(venv.join("pyvenv.cfg"), "home = /usr/bin\n").unwrap();
|
||||
fs::write(venv.join("unsloth_install_manifest.json"), "{}").unwrap();
|
||||
|
||||
let site_packages = venv.join("lib").join("python3.11").join("site-packages");
|
||||
// site_packages_dirs() only walks lib/<pyver>/site-packages on unix; on
|
||||
// Windows it looks at Lib/site-packages. Building the posix layout
|
||||
// everywhere left the dist-info invisible to the fingerprint on Windows,
|
||||
// so removing it changed nothing and the assert_ne below could not hold.
|
||||
let site_packages = if cfg!(windows) {
|
||||
venv.join("Lib").join("site-packages")
|
||||
} else {
|
||||
venv.join("lib").join("python3.11").join("site-packages")
|
||||
};
|
||||
fs::create_dir_all(site_packages.join("unsloth_cli").join("commands")).unwrap();
|
||||
fs::write(
|
||||
site_packages
|
||||
|
|
|
|||
|
|
@ -35,6 +35,11 @@
|
|||
]
|
||||
},
|
||||
"plugins": {
|
||||
"deep-link": {
|
||||
"desktop": {
|
||||
"schemes": ["unsloth"]
|
||||
}
|
||||
},
|
||||
"updater": {
|
||||
"pubkey": "dW50cnVzdGVkIGNvbW1lbnQ6IG1pbmlzaWduIHB1YmxpYyBrZXk6IDFBQzA4RjczODM0RjE1QjcKUldTM0ZVK0RjNC9BR2t4R0RVaFR5cTkyUlRVQ1FwaGV0Nk04eWNwWXBhZnlzalJydllmZm1QTS8K",
|
||||
"endpoints": [
|
||||
|
|
@ -80,6 +85,7 @@
|
|||
"bundleMediaFramework": false
|
||||
},
|
||||
"deb": {
|
||||
"desktopTemplate": "./linux/unsloth.desktop",
|
||||
"postRemoveScript": "./linux/postremove.sh"
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -862,3 +862,76 @@ class TestNoTorchPersistenceParity:
|
|||
manifest = (REPO_ROOT / "studio" / "install_manifest.py").read_text(encoding = "utf-8")
|
||||
assert 'NO_TORCH_TRUTHY: Tuple[str, ...] = ("1", "true", "yes", "on")' in manifest
|
||||
assert "install_manifest.NO_TORCH_TRUTHY" in STACK_PY.read_text(encoding = "utf-8")
|
||||
|
||||
|
||||
class TestAmdBnbFloorParity:
|
||||
"""bitsandbytes <= 0.49.2 NaNs at 4-bit decode shape on every AMD GPU; the ROCm
|
||||
4-bit GEMV fix (bnb #1887) first ships on PyPI in 0.50.0. The `amd` extra,
|
||||
install.sh and the Studio stack resolve bitsandbytes independently, so all three
|
||||
must carry the same floor or an unreachable pre-release wheel silently reinstates
|
||||
the broken range."""
|
||||
|
||||
FLOOR = "0.50.0"
|
||||
PYPROJECT = REPO_ROOT / "pyproject.toml"
|
||||
|
||||
def test_amd_extra_floor(self):
|
||||
text = self.PYPROJECT.read_text(encoding = "utf-8")
|
||||
amd = re.search(r"^amd = \[(.*?)^\]", text, re.S | re.M)
|
||||
assert amd, "pyproject.toml must define an `amd` extra"
|
||||
specs = re.findall(r'"(bitsandbytes[^"]*)"', amd.group(1))
|
||||
assert specs, "the amd extra must pin bitsandbytes"
|
||||
for spec in specs:
|
||||
assert spec.startswith(
|
||||
f"bitsandbytes>={self.FLOOR}"
|
||||
), f"amd extra bitsandbytes floor must be >={self.FLOOR}, got {spec!r}"
|
||||
|
||||
def test_install_sh_pypi_fallback_floor(self):
|
||||
text = INSTALL_SH.read_text(encoding = "utf-8")
|
||||
assert (
|
||||
f'_BNB_ROCM_PYPI_FALLBACK="bitsandbytes>={self.FLOOR}"' in text
|
||||
), f"install.sh _install_bnb_rocm PyPI fallback must floor at {self.FLOOR}"
|
||||
|
||||
def test_stack_py_pypi_fallback_floor(self):
|
||||
text = STACK_PY.read_text(encoding = "utf-8")
|
||||
assert (
|
||||
f'_BNB_ROCM_PYPI_FALLBACK = "bitsandbytes>={self.FLOOR}"' in text
|
||||
), f"install_python_stack.py PyPI fallback must floor at {self.FLOOR}"
|
||||
|
||||
def test_no_installer_still_allows_the_broken_range(self):
|
||||
for path in (INSTALL_SH, INSTALL_PS1, SETUP_PS1, STACK_PY, self.PYPROJECT):
|
||||
text = path.read_text(encoding = "utf-8")
|
||||
for line in text.splitlines():
|
||||
if "bitsandbytes>=0.49" in line and not line.lstrip().startswith(("#", "//")):
|
||||
raise AssertionError(
|
||||
f"{path.name} still floors bitsandbytes in the broken ROCm range: {line.strip()!r}"
|
||||
)
|
||||
|
||||
def test_fallback_is_not_reported_as_broken(self):
|
||||
"""The fallback now installs the first fixed release, so neither installer
|
||||
may still call 4-bit decode broken on ROCm."""
|
||||
for path in (INSTALL_SH, STACK_PY):
|
||||
text = path.read_text(encoding = "utf-8")
|
||||
assert (
|
||||
"4-bit decode broken on ROCm" not in text
|
||||
), f"{path.name} still reports the repaired PyPI fallback as broken"
|
||||
assert (
|
||||
"4-bit decode will be broken on ROCm" not in text
|
||||
), f"{path.name} still reports the repaired PyPI fallback as broken"
|
||||
|
||||
def test_aarch64_is_not_told_it_has_a_rocm_backend(self):
|
||||
"""bitsandbytes ships no ROCm kernels in its aarch64 wheel at any version, so
|
||||
neither installer may hand aarch64 the x86_64 "carries the ROCm 4-bit fix"
|
||||
message, and both must warn that 4-bit needs a source build there."""
|
||||
sh = INSTALL_SH.read_text(encoding = "utf-8")
|
||||
assert "_bnb_rocm_arch_has_binary()" in sh
|
||||
assert "_warn_bnb_no_rocm_binary()" in sh
|
||||
assert (
|
||||
sh.count("_warn_bnb_no_rocm_binary\n") >= 2
|
||||
), "install.sh must warn on aarch64 after both the pre-release and the fallback install"
|
||||
py = STACK_PY.read_text(encoding = "utf-8")
|
||||
assert "def _bnb_rocm_arch_has_binary(" in py
|
||||
assert "_bnb_rocm_arch_has_binary()" in py
|
||||
for text, name in ((sh, "install.sh"), (py, "install_python_stack.py")):
|
||||
assert (
|
||||
"4-bit QLoRA needs a source build" in text
|
||||
), f"{name} must tell aarch64 users 4-bit needs a source build"
|
||||
|
|
|
|||
296
tests/python/test_import_without_bitsandbytes.py
Normal file
296
tests/python/test_import_without_bitsandbytes.py
Normal file
|
|
@ -0,0 +1,296 @@
|
|||
"""`import unsloth` must survive a missing bitsandbytes.
|
||||
|
||||
device_type.py already tells the user "bitsandbytes is not installed - 4bit QLoRA
|
||||
unallowed, but 16bit and full finetuning works", and the gfx906 install path
|
||||
(#7354) deliberately removes the generic wheel because it carries no gfx906
|
||||
kernels. Any module-level `import bitsandbytes` on the import chain turns that
|
||||
into an unimportable package instead.
|
||||
|
||||
peft's 4bit LoRA layer is exported only when bnb is importable, so
|
||||
`from peft.tuners.lora import Linear4bit` fails on the same hosts and is checked
|
||||
here too.
|
||||
"""
|
||||
|
||||
# Path | None below is a PEP 604 union; the project still supports Python 3.9.
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
ROOT_MODULE = "unsloth"
|
||||
|
||||
|
||||
def _module_path(name: str) -> Path | None:
|
||||
base = REPO_ROOT / Path(*name.split("."))
|
||||
for candidate in (base.with_suffix(".py"), base / "__init__.py"):
|
||||
if candidate.is_file():
|
||||
return candidate
|
||||
return None
|
||||
|
||||
|
||||
def _bnb_dependent(node: ast.stmt) -> bool:
|
||||
"""True for an import that raises when bitsandbytes is absent."""
|
||||
if isinstance(node, ast.Import):
|
||||
return any(a.name.split(".")[0] == "bitsandbytes" for a in node.names)
|
||||
if isinstance(node, ast.ImportFrom) and node.level == 0:
|
||||
module = node.module or ""
|
||||
if module.split(".")[0] == "bitsandbytes":
|
||||
return True
|
||||
# peft re-exports Linear4bit only when bnb imported cleanly.
|
||||
if module.startswith("peft.tuners.lora"):
|
||||
return any(a.name == "Linear4bit" for a in node.names)
|
||||
return False
|
||||
|
||||
|
||||
def _allow_bitsandbytes_gated(test: ast.expr) -> bool:
|
||||
"""device_type.py sets ALLOW_BITSANDBYTES=False exactly when the import failed,
|
||||
so a branch keyed on it cannot run without bnb."""
|
||||
return any(isinstance(n, ast.Name) and n.id == "ALLOW_BITSANDBYTES" for n in ast.walk(test))
|
||||
|
||||
|
||||
def _scan(path: Path, module: str):
|
||||
"""Yield (lineno, source) for unguarded top-level imports.
|
||||
|
||||
Imports inside a `try`, or under an ALLOW_BITSANDBYTES branch, are guarded.
|
||||
Other `if` bodies are not: the condition may well be true on a host without bnb.
|
||||
"""
|
||||
is_package = path.name == "__init__.py"
|
||||
package = module if is_package else module.rpartition(".")[0]
|
||||
tree = ast.parse(path.read_text(encoding = "utf-8"))
|
||||
risky, edges = [], []
|
||||
|
||||
def walk(body, guarded):
|
||||
for node in body:
|
||||
if isinstance(node, (ast.Import, ast.ImportFrom)):
|
||||
if not guarded and _bnb_dependent(node):
|
||||
risky.append((node.lineno, ast.unparse(node)))
|
||||
if isinstance(node, ast.Import):
|
||||
edges.extend(a.name for a in node.names)
|
||||
elif node.level:
|
||||
parts = package.split(".")
|
||||
base = ".".join(parts[: len(parts) - (node.level - 1)])
|
||||
edges.append(f"{base}.{node.module}" if node.module else base)
|
||||
else:
|
||||
edges.append(node.module or "")
|
||||
elif isinstance(node, ast.Try):
|
||||
walk(node.body, True)
|
||||
for handler in node.handlers:
|
||||
walk(handler.body, True)
|
||||
walk(node.orelse, True)
|
||||
walk(node.finalbody, guarded)
|
||||
elif isinstance(node, ast.If):
|
||||
walk(node.body, guarded or _allow_bitsandbytes_gated(node.test))
|
||||
walk(node.orelse, guarded)
|
||||
|
||||
walk(tree.body, False)
|
||||
return risky, edges
|
||||
|
||||
|
||||
def test_no_unguarded_bitsandbytes_import_on_the_unsloth_import_chain():
|
||||
seen, pending, offenders = set(), [(ROOT_MODULE, [])], []
|
||||
while pending:
|
||||
module, chain = pending.pop()
|
||||
if module in seen:
|
||||
continue
|
||||
seen.add(module)
|
||||
path = _module_path(module)
|
||||
if path is None:
|
||||
continue
|
||||
risky, edges = _scan(path, module)
|
||||
for lineno, source in risky:
|
||||
rel = path.relative_to(REPO_ROOT).as_posix()
|
||||
offenders.append(f"{rel}:{lineno} {source}\n via {' -> '.join(chain + [module])}")
|
||||
pending.extend(
|
||||
(edge, chain + [module]) for edge in edges if edge.split(".")[0] == ROOT_MODULE
|
||||
)
|
||||
|
||||
assert len(seen) > 20, f"import chain walk collapsed, only reached {seen}"
|
||||
assert not offenders, (
|
||||
"`import unsloth` must not hard-require bitsandbytes. Wrap these in "
|
||||
"try/except and fall back to a placeholder:\n " + "\n ".join(offenders)
|
||||
)
|
||||
|
||||
|
||||
def test_missing_bnb_leaves_a_callable_that_reports_the_real_cause():
|
||||
"""The 4bit ctypes handles degrade to a stub, not a NameError later on."""
|
||||
src = (REPO_ROOT / "unsloth" / "kernels" / "utils.py").read_text(encoding = "utf-8")
|
||||
assert "def _bnb_required(" in src
|
||||
assert "get_ptr = _bnb_required" in src
|
||||
for name in (
|
||||
"cdequantize_blockwise_fp32",
|
||||
"cdequantize_blockwise_fp16_nf4",
|
||||
"cdequantize_blockwise_bf16_nf4",
|
||||
"cgemm_4bit_inference_naive_fp16",
|
||||
"cgemm_4bit_inference_naive_bf16",
|
||||
):
|
||||
assert f"{name} = _bnb_required" in src, f"{name} has no bnb-less fallback"
|
||||
|
||||
|
||||
def test_capability_flags_come_from_a_guarded_import_not_find_spec():
|
||||
"""kernels/utils.py and _gpu_init.py treat any import failure as unavailable.
|
||||
device_type.py must agree, or an installed-but-unusable wheel leaves
|
||||
ALLOW_BITSANDBYTES true while the kernels fall back to the stub."""
|
||||
src = (REPO_ROOT / "unsloth" / "device_type.py").read_text(encoding = "utf-8")
|
||||
head = src.split('if DEVICE_TYPE == "hip":')[0]
|
||||
assert "import bitsandbytes as _bnb_probe" in head
|
||||
assert 'find_spec("bitsandbytes")' not in head, "find_spec cannot see a broken wheel"
|
||||
assert head.count("ALLOW_BITSANDBYTES = False") >= 1
|
||||
|
||||
|
||||
def _bnb_guards():
|
||||
src = (REPO_ROOT / "unsloth" / "models" / "loader.py").read_text(encoding = "utf-8")
|
||||
tree = ast.parse(src)
|
||||
return src, [
|
||||
node
|
||||
for node in ast.walk(tree)
|
||||
if isinstance(node, ast.If)
|
||||
and any(
|
||||
isinstance(n, ast.Name) and n.id == "ALLOW_BITSANDBYTES" for n in ast.walk(node.test)
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
def test_bitsandbytes_guard_is_not_gated_on_use_exact_model_name():
|
||||
"""use_exact_model_name suppresses repo-name remapping; it cannot make bnb
|
||||
available. Gating on it left the default load_in_4bit=True set on a host
|
||||
without bitsandbytes."""
|
||||
_, guards = _bnb_guards()
|
||||
assert len(guards) == 2, f"expected both loader guards, found {len(guards)}"
|
||||
for guard in guards:
|
||||
names = {n.id for n in ast.walk(guard.test) if isinstance(n, ast.Name)}
|
||||
assert (
|
||||
"use_exact_model_name" not in names
|
||||
), f"guard at line {guard.lineno} still gates the capability check on naming"
|
||||
|
||||
|
||||
def test_bitsandbytes_guard_drops_a_bnb_quantization_config():
|
||||
"""A BitsAndBytesConfig in kwargs re-sets the flags downstream, so clearing
|
||||
load_in_4bit/8bit alone still builds the bnb quantizer in Transformers. A
|
||||
non-bnb config (GPTQ/AWQ/fp8) must not be touched."""
|
||||
_, guards = _bnb_guards()
|
||||
for guard in guards:
|
||||
# ast.unparse normalises quotes, so match on the call shape instead.
|
||||
def _is_pop(node):
|
||||
return (
|
||||
isinstance(node, ast.Call)
|
||||
and isinstance(node.func, ast.Attribute)
|
||||
and node.func.attr == "pop"
|
||||
and isinstance(node.func.value, ast.Name)
|
||||
and node.func.value.id == "kwargs"
|
||||
and node.args
|
||||
and isinstance(node.args[0], ast.Constant)
|
||||
and node.args[0].value == "quantization_config"
|
||||
)
|
||||
|
||||
assert any(
|
||||
_is_pop(n) for n in ast.walk(guard)
|
||||
), f"guard at line {guard.lineno} leaves the bnb config in kwargs"
|
||||
# the pop must be conditional on the config actually asking for bnb
|
||||
pops = [
|
||||
node
|
||||
for node in ast.walk(guard)
|
||||
if isinstance(node, ast.If) and any(_is_pop(n) for n in ast.walk(node))
|
||||
]
|
||||
assert pops, f"guard at line {guard.lineno} pops unconditionally"
|
||||
assert any(
|
||||
isinstance(n, ast.Name) and n.id == "_wants_bnb"
|
||||
for node in pops
|
||||
for n in ast.walk(node.test)
|
||||
), f"guard at line {guard.lineno} does not gate the pop on a bnb request"
|
||||
|
||||
|
||||
def test_bitsandbytes_guard_clears_8bit_as_well_as_4bit():
|
||||
"""8bit is bitsandbytes too: leaving load_in_8bit set sends the request to
|
||||
Transformers, which builds the bnb quantizer and fails there instead."""
|
||||
src = (REPO_ROOT / "unsloth" / "models" / "loader.py").read_text(encoding = "utf-8")
|
||||
tree = ast.parse(src)
|
||||
guards = [
|
||||
node
|
||||
for node in ast.walk(tree)
|
||||
if isinstance(node, ast.If)
|
||||
and any(
|
||||
isinstance(n, ast.Name) and n.id == "ALLOW_BITSANDBYTES" for n in ast.walk(node.test)
|
||||
)
|
||||
]
|
||||
assert len(guards) == 2, f"expected both loader guards, found {len(guards)}"
|
||||
for guard in guards:
|
||||
cleared = {
|
||||
target.id
|
||||
for stmt in guard.body
|
||||
if isinstance(stmt, ast.Assign)
|
||||
for target in stmt.targets
|
||||
if isinstance(target, ast.Name)
|
||||
and isinstance(stmt.value, ast.Constant)
|
||||
and stmt.value.value is False
|
||||
}
|
||||
assert {
|
||||
"load_in_4bit",
|
||||
"load_in_8bit",
|
||||
} <= cleared, f"guard at line {guard.lineno} clears only {sorted(cleared)}"
|
||||
|
||||
|
||||
def test_capability_fallback_precedes_the_mutually_exclusive_mode_check():
|
||||
"""load_in_4bit defaults to True, so load_in_16bit=True trips the
|
||||
"can only load in 4bit or 8bit or 16bit" RuntimeError unless the unavailable
|
||||
4bit request is cleared first. That check must come after the fallback."""
|
||||
src, _ = _bnb_guards()
|
||||
tree = ast.parse(src)
|
||||
checked = 0
|
||||
# Scope to the enclosing function: the other loader's guard sits earlier in the
|
||||
# file and would otherwise satisfy a plain line-number comparison.
|
||||
for func in ast.walk(tree):
|
||||
if not isinstance(func, ast.FunctionDef):
|
||||
continue
|
||||
raises = [
|
||||
node.lineno
|
||||
for node in ast.walk(func)
|
||||
if isinstance(node, ast.Raise)
|
||||
and "Can only load in 4bit or 8bit or 16bit" in ast.unparse(node)
|
||||
]
|
||||
if not raises:
|
||||
continue
|
||||
guards = [
|
||||
node.lineno
|
||||
for node in ast.walk(func)
|
||||
if isinstance(node, ast.If)
|
||||
and any(
|
||||
isinstance(n, ast.Name) and n.id == "ALLOW_BITSANDBYTES"
|
||||
for n in ast.walk(node.test)
|
||||
)
|
||||
]
|
||||
for lineno in raises:
|
||||
checked += 1
|
||||
assert any(g < lineno for g in guards), (
|
||||
f"{func.name}: the mode check at line {lineno} runs before this "
|
||||
"function's ALLOW_BITSANDBYTES fallback, so load_in_16bit=True on a "
|
||||
"bnb-less host raises instead of taking the 16bit path"
|
||||
)
|
||||
assert checked, "mode-exclusivity check not found"
|
||||
|
||||
|
||||
def test_bitsandbytes_compile_patch_is_never_called_unguarded():
|
||||
"""unsloth_zoo's patch_compiling_bitsandbytes imports bitsandbytes
|
||||
unconditionally, so an unwrapped call raises on a bnb-less host before any
|
||||
fallback can run."""
|
||||
src = (REPO_ROOT / "unsloth" / "models" / "loader.py").read_text(encoding = "utf-8")
|
||||
tree = ast.parse(src)
|
||||
calls = [
|
||||
node
|
||||
for node in ast.walk(tree)
|
||||
if isinstance(node, ast.Call)
|
||||
and isinstance(node.func, ast.Name)
|
||||
and node.func.id == "patch_compiling_bitsandbytes"
|
||||
]
|
||||
assert calls, "call sites not found"
|
||||
guarded = {
|
||||
call.lineno
|
||||
for node in ast.walk(tree)
|
||||
if isinstance(node, ast.Try)
|
||||
for call in ast.walk(node)
|
||||
if isinstance(call, ast.Call)
|
||||
and isinstance(call.func, ast.Name)
|
||||
and call.func.id == "patch_compiling_bitsandbytes"
|
||||
}
|
||||
unguarded = sorted({c.lineno for c in calls} - guarded)
|
||||
assert not unguarded, f"patch_compiling_bitsandbytes called unguarded at {unguarded}"
|
||||
|
|
@ -96,12 +96,11 @@ def load_and_compute_8bit_ppl(
|
|||
if __name__ == "__main__":
|
||||
mp.set_start_method("spawn", force = True)
|
||||
|
||||
if torch.cuda.is_bf16_supported():
|
||||
compute_dtype = torch.bfloat16
|
||||
attn_implementation = "flash_attention_2"
|
||||
else:
|
||||
compute_dtype = torch.float16
|
||||
attn_implementation = "sdpa"
|
||||
from unsloth import is_bfloat16_supported
|
||||
from unsloth.models._utils import HAS_FLASH_ATTENTION
|
||||
|
||||
compute_dtype = torch.bfloat16 if is_bfloat16_supported() else torch.float16
|
||||
attn_implementation = "flash_attention_2" if HAS_FLASH_ATTENTION else "sdpa"
|
||||
|
||||
model, tokenizer = FastLanguageModel.from_pretrained(
|
||||
model_name = "unsloth/Llama-3.2-3B-Instruct",
|
||||
|
|
|
|||
|
|
@ -121,12 +121,11 @@ def load_and_compute_8bit_ppl(
|
|||
if __name__ == "__main__":
|
||||
mp.set_start_method("spawn", force = True)
|
||||
|
||||
if torch.cuda.is_bf16_supported():
|
||||
compute_dtype = torch.bfloat16
|
||||
attn_implementation = "flash_attention_2"
|
||||
else:
|
||||
compute_dtype = torch.float16
|
||||
attn_implementation = "sdpa"
|
||||
from unsloth import is_bfloat16_supported
|
||||
from unsloth.models._utils import HAS_FLASH_ATTENTION
|
||||
|
||||
compute_dtype = torch.bfloat16 if is_bfloat16_supported() else torch.float16
|
||||
attn_implementation = "flash_attention_2" if HAS_FLASH_ATTENTION else "sdpa"
|
||||
|
||||
model, tokenizer = FastLanguageModel.from_pretrained(
|
||||
model_name = "unsloth/mistral-7b-v0.3",
|
||||
|
|
|
|||
|
|
@ -98,12 +98,11 @@ def load_and_compute_8bit_ppl(
|
|||
if __name__ == "__main__":
|
||||
mp.set_start_method("spawn", force = True)
|
||||
|
||||
if torch.cuda.is_bf16_supported():
|
||||
compute_dtype = torch.bfloat16
|
||||
attn_implementation = "flash_attention_2"
|
||||
else:
|
||||
compute_dtype = torch.float16
|
||||
attn_implementation = "sdpa"
|
||||
from unsloth import is_bfloat16_supported
|
||||
from unsloth.models._utils import HAS_FLASH_ATTENTION
|
||||
|
||||
compute_dtype = torch.bfloat16 if is_bfloat16_supported() else torch.float16
|
||||
attn_implementation = "flash_attention_2" if HAS_FLASH_ATTENTION else "sdpa"
|
||||
|
||||
model, tokenizer = FastLanguageModel.from_pretrained(
|
||||
model_name = "unsloth/Phi-4",
|
||||
|
|
|
|||
|
|
@ -95,12 +95,11 @@ def load_and_compute_8bit_ppl(
|
|||
if __name__ == "__main__":
|
||||
mp.set_start_method("spawn", force = True)
|
||||
|
||||
if torch.cuda.is_bf16_supported():
|
||||
compute_dtype = torch.bfloat16
|
||||
attn_implementation = "flash_attention_2"
|
||||
else:
|
||||
compute_dtype = torch.float16
|
||||
attn_implementation = "sdpa"
|
||||
from unsloth import is_bfloat16_supported
|
||||
from unsloth.models._utils import HAS_FLASH_ATTENTION
|
||||
|
||||
compute_dtype = torch.bfloat16 if is_bfloat16_supported() else torch.float16
|
||||
attn_implementation = "flash_attention_2" if HAS_FLASH_ATTENTION else "sdpa"
|
||||
|
||||
model, tokenizer = FastLanguageModel.from_pretrained(
|
||||
model_name = "unsloth/Llama-3.1-8B-Instruct",
|
||||
|
|
|
|||
|
|
@ -164,12 +164,11 @@ def load_and_compute_8bit_ppl(
|
|||
if __name__ == "__main__":
|
||||
mp.set_start_method("spawn", force = True)
|
||||
|
||||
if torch.cuda.is_bf16_supported():
|
||||
compute_dtype = torch.bfloat16
|
||||
attn_implementation = "flash_attention_2"
|
||||
else:
|
||||
compute_dtype = torch.float16
|
||||
attn_implementation = "sdpa"
|
||||
from unsloth import is_bfloat16_supported
|
||||
from unsloth.models._utils import HAS_FLASH_ATTENTION
|
||||
|
||||
compute_dtype = torch.bfloat16 if is_bfloat16_supported() else torch.float16
|
||||
attn_implementation = "flash_attention_2" if HAS_FLASH_ATTENTION else "sdpa"
|
||||
|
||||
model, tokenizer = FastLanguageModel.from_pretrained(
|
||||
model_name = "unsloth/Qwen2.5-7B-Instruct",
|
||||
|
|
@ -210,8 +209,6 @@ if __name__ == "__main__":
|
|||
loftq_config = None,
|
||||
)
|
||||
|
||||
from unsloth import is_bfloat16_supported
|
||||
|
||||
trainer = SFTTrainer(
|
||||
model = model,
|
||||
tokenizer = tokenizer,
|
||||
|
|
|
|||
|
|
@ -124,7 +124,12 @@ def _call(
|
|||
tmp_path factory branches on it, so a session-wide patch breaks the fixture on
|
||||
a Windows test host."""
|
||||
with patch.object(sys, "platform", platform):
|
||||
with patch("os.path.exists", _fake_exists(present)):
|
||||
# isdir too: the llvm probe requires a directory, so a fake host that only
|
||||
# answers exists() would report every nested llvm dir as missing.
|
||||
with (
|
||||
patch("os.path.exists", _fake_exists(present)),
|
||||
patch("os.path.isdir", _fake_exists(present)),
|
||||
):
|
||||
return _norm(impl(str(bundle)))
|
||||
|
||||
|
||||
|
|
@ -260,6 +265,78 @@ class TestNativeLinuxRootResolution:
|
|||
for where, impl in _impls().items():
|
||||
assert self._run(impl, bundle_dir, present) == ["/opt/rocm/lib64"], where
|
||||
|
||||
def test_nested_llvm_runtime_follows_system_rocm_lib(self, bundle_dir):
|
||||
"""#7446: libamd_comgr depends on ROCm's versioned LLVM runtime, which is
|
||||
installed below lib/llvm/lib rather than directly in lib."""
|
||||
present = {
|
||||
"/dev/kfd",
|
||||
"/opt/rocm/lib/libhsa-runtime64.so",
|
||||
"/opt/rocm/lib/llvm/lib",
|
||||
}
|
||||
for where, impl in _impls().items():
|
||||
assert self._run(impl, bundle_dir, present) == [
|
||||
"/opt/rocm/lib",
|
||||
"/opt/rocm/lib/llvm/lib",
|
||||
], where
|
||||
|
||||
def test_lib64_host_still_finds_llvm_under_lib(self, bundle_dir):
|
||||
"""ROCm puts LLVM under <root>/lib/llvm even where HSA lives in lib64, so
|
||||
deriving the llvm dir from the HSA dir alone would miss it and leave
|
||||
libamd_comgr binding to the bundle's libLLVM."""
|
||||
present = {
|
||||
"/dev/kfd",
|
||||
"/opt/rocm/lib64/libhsa-runtime64.so",
|
||||
"/opt/rocm/lib/llvm/lib",
|
||||
}
|
||||
for where, impl in _impls().items():
|
||||
assert self._run(impl, bundle_dir, present) == [
|
||||
"/opt/rocm/lib64",
|
||||
"/opt/rocm/lib/llvm/lib",
|
||||
], where
|
||||
|
||||
def test_lib64_host_prefers_its_own_llvm_dir_when_both_exist(self, bundle_dir):
|
||||
present = {
|
||||
"/dev/kfd",
|
||||
"/opt/rocm/lib64/libhsa-runtime64.so",
|
||||
"/opt/rocm/lib64/llvm/lib",
|
||||
"/opt/rocm/lib/llvm/lib",
|
||||
}
|
||||
for where, impl in _impls().items():
|
||||
assert self._run(impl, bundle_dir, present) == [
|
||||
"/opt/rocm/lib64",
|
||||
"/opt/rocm/lib64/llvm/lib",
|
||||
"/opt/rocm/lib/llvm/lib",
|
||||
], where
|
||||
|
||||
def test_llvm_path_that_is_a_file_is_not_prepended(self, tmp_path, bundle_dir):
|
||||
"""Real filesystem: the serve-time caller joins these straight into
|
||||
LD_LIBRARY_PATH without an is-dir filter, so a non-directory must not
|
||||
reach it."""
|
||||
root = tmp_path / "rocm"
|
||||
(root / "lib").mkdir(parents = True)
|
||||
(root / "lib" / "libhsa-runtime64.so").write_text("")
|
||||
(root / "lib" / "llvm").mkdir()
|
||||
(root / "lib" / "llvm" / "lib").write_text("not a directory")
|
||||
real_exists = os.path.exists
|
||||
# Pin both device nodes: a WSL test host really has /dev/dxg, which would
|
||||
# take the WSL early-return and make this pass for the wrong reason.
|
||||
pinned = {"/dev/kfd": True, "/dev/dxg": False}
|
||||
|
||||
def _exists(p):
|
||||
return pinned.get(str(p), None) if str(p) in pinned else real_exists(p)
|
||||
|
||||
# A test host may itself have a real /opt/rocm (the default candidate), so
|
||||
# assert on the bogus entry rather than on the whole list.
|
||||
for where, impl in _impls().items():
|
||||
with (
|
||||
patch.object(sys, "platform", "linux"),
|
||||
patch.dict(os.environ, {"ROCM_PATH": str(root)}, clear = True),
|
||||
patch("os.path.exists", _exists),
|
||||
):
|
||||
out = impl(str(bundle_dir))
|
||||
assert str(root / "lib") in out, where
|
||||
assert str(root / "lib" / "llvm" / "lib") not in out, where
|
||||
|
||||
def test_lib_precedes_lib64_when_both_exist(self, bundle_dir):
|
||||
present = {
|
||||
"/dev/kfd",
|
||||
|
|
|
|||
|
|
@ -3157,12 +3157,35 @@ class TestInstallBnbWindowsRocm:
|
|||
assert result is False
|
||||
assert "BNB_ROCM_VERSION" not in os.environ
|
||||
|
||||
def test_no_op_when_win_amd64_url_missing(self):
|
||||
"""Should be silent no-op if win_amd64 key absent from _BNB_ROCM_PRERELEASE_URLS."""
|
||||
def test_falls_back_to_pypi_when_win_amd64_url_missing(self):
|
||||
"""No win_amd64 pre-release wheel must not mean no bitsandbytes: PyPI
|
||||
>=0.50.0 ships libbitsandbytes_rocm{714,72}.dll, so it is a real ROCm build."""
|
||||
with patch.object(stack_mod, "_BNB_ROCM_PRERELEASE_URLS", {}):
|
||||
with patch.object(stack_mod, "pip_install_try") as mock_pip:
|
||||
with patch.object(stack_mod, "pip_install_try", return_value = True) as mock_pip:
|
||||
stack_mod._install_bnb_windows_rocm()
|
||||
mock_pip.assert_not_called()
|
||||
assert mock_pip.call_count == 1
|
||||
assert stack_mod._BNB_ROCM_PYPI_FALLBACK in mock_pip.call_args.args
|
||||
|
||||
def test_falls_back_to_pypi_when_prerelease_install_fails(self):
|
||||
"""A blocked GitHub pre-release URL must fall through to the PyPI floor rather
|
||||
than leaving Windows ROCm with no working bitsandbytes."""
|
||||
with patch.object(stack_mod, "pip_install_try", side_effect = [False, True]) as mock_pip:
|
||||
with patch.object(stack_mod, "_detect_bnb_rocm_dll_ver", return_value = "72"):
|
||||
result = stack_mod._install_bnb_windows_rocm()
|
||||
assert result is True
|
||||
assert mock_pip.call_count == 2
|
||||
assert "win_amd64" in str(mock_pip.call_args_list[0])
|
||||
assert stack_mod._BNB_ROCM_PYPI_FALLBACK in mock_pip.call_args_list[1].args
|
||||
|
||||
def test_returns_false_only_when_both_paths_fail(self):
|
||||
"""Both the pre-release wheel and the PyPI fallback must fail before the
|
||||
helper reports failure."""
|
||||
with patch.dict(os.environ, {}, clear = False):
|
||||
os.environ.pop("BNB_ROCM_VERSION", None)
|
||||
with patch.object(stack_mod, "pip_install_try", return_value = False) as mock_pip:
|
||||
result = stack_mod._install_bnb_windows_rocm()
|
||||
assert result is False
|
||||
assert mock_pip.call_count == 2
|
||||
|
||||
def test_sets_bnb_rocm_version_from_detected_dll(self):
|
||||
"""BNB_ROCM_VERSION is set from the DLL detected after install."""
|
||||
|
|
|
|||
|
|
@ -68,3 +68,18 @@ def test_backend_chat_preset_accepts_load_config():
|
|||
routes = _read("studio/backend/routes/chat_history.py")
|
||||
assert "class ChatPresetLoadConfig" in routes
|
||||
assert "loadConfig: Optional[ChatPresetLoadConfig]" in routes
|
||||
|
||||
|
||||
def test_preset_load_config_carries_parallel_slots():
|
||||
# Captured, clamped on read, applied, and accepted by the extra="forbid"
|
||||
# backend model (a missing backend field would 422 every settings sync).
|
||||
source = _read("studio/frontend/src/features/chat/presets/preset-load-config.ts")
|
||||
assert '| "nParallel"' in source
|
||||
assert "nParallel: snapshot.nParallel ?? null" in source
|
||||
assert "nParallel: config.nParallel ?? null" in source
|
||||
assert "N_PARALLEL_MAX, Math.round(partial.nParallel)" in source
|
||||
routes = _read("studio/backend/routes/chat_history.py")
|
||||
assert (
|
||||
"nParallel: Optional[int] = Field(default = None, ge = PARALLEL_MIN, le = PARALLEL_MAX)"
|
||||
in routes
|
||||
)
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ APP_PROVIDER = FRONTEND / "app/provider.tsx"
|
|||
|
||||
CLIPBOARD_FILES = FRONTEND / "features/chat/utils/clipboard-files.ts"
|
||||
TAURI_CAPABILITIES = REPO / "studio/src-tauri/capabilities/default.json"
|
||||
CHAT_PAGE = FRONTEND / "features/chat/chat-page.tsx"
|
||||
|
||||
|
||||
def test_file_actions_route_through_native_commands_only_in_tauri():
|
||||
|
|
@ -216,6 +217,34 @@ def test_expanded_titlebar_button_and_corner_match_sidebar_edge():
|
|||
)
|
||||
|
||||
|
||||
def test_visible_mac_sidebar_header_is_a_drag_region():
|
||||
source = APP_SIDEBAR.read_text(encoding = "utf-8")
|
||||
header = source.split("<SidebarHeader", 1)[1].split("</SidebarHeader>", 1)[0]
|
||||
drag_region = "data-tauri-drag-region={usesNativeMacTitlebar || undefined}"
|
||||
|
||||
assert drag_region in header
|
||||
assert header.index(drag_region) < header.index('"relative z-10 flex items-center')
|
||||
|
||||
|
||||
def test_mac_chat_header_controls_share_the_titlebar_row():
|
||||
source = CHAT_PAGE.read_text(encoding = "utf-8")
|
||||
provider = APP_PROVIDER.read_text(encoding = "utf-8")
|
||||
|
||||
assert "shouldUseNativeMacWindowTitlebar" not in source
|
||||
assert "[--studio-content-top-inset:var(--studio-mac-titlebar-height" not in source
|
||||
assert source.count("var(--studio-mac-traffic-light-inset") == 2
|
||||
assert '"--studio-chat-header-padding-top": "7px"' in provider
|
||||
assert "pt-[var(--studio-content-top-inset,0px)] md:flex-row" in source
|
||||
assert "absolute top-[var(--studio-content-top-inset,0px)]" in source
|
||||
|
||||
|
||||
def test_collapsed_mac_sidebar_hides_divider():
|
||||
source = APP_SIDEBAR.read_text(encoding = "utf-8")
|
||||
|
||||
assert "group-data-[collapsible=icon]:[&_[data-sidebar=sidebar]]:border-r-0" in source
|
||||
assert "top-[var(--studio-mac-titlebar-height,34px)]" not in source
|
||||
|
||||
|
||||
def test_chat_sidebar_row_actions_visible_on_coarse_pointers():
|
||||
"""unslothai/unsloth#7276: Recents chat kebab must be tappable on iPad."""
|
||||
sidebar_source = APP_SIDEBAR.read_text(encoding = "utf-8")
|
||||
|
|
|
|||
|
|
@ -642,6 +642,253 @@ def test_legacy_migration_is_idempotent_and_non_destructive():
|
|||
assert "if (isDefaultConfig(migrated) || Object.hasOwn(map, key)) {" in src
|
||||
|
||||
|
||||
def test_parallel_slots_setting_wired_end_to_end():
|
||||
"""The per-load Parallel Slots knob (llama-server --parallel) must flow from
|
||||
the run-settings form through persistence, every /load builder, the validate
|
||||
preflight and the cross-model reset; a lost hop silently reverts the model to
|
||||
the server-wide slot default."""
|
||||
config = _read("features/model-picker/model-config/per-model-config.ts")
|
||||
# Persisted per model, clamped on every read/write, and null (= server
|
||||
# default) counts as default so blank configs are not stored.
|
||||
assert '"nParallel",' in config
|
||||
assert "N_PARALLEL_MAX, Math.round(partial.nParallel)" in config
|
||||
assert "config.nParallel == null &&" in config
|
||||
page = _read("features/model-picker/components/model-config-page.tsx")
|
||||
# Rendered in the GGUF advanced section, which a remembered override reopens.
|
||||
assert "Parallel Slots" in page
|
||||
assert "config.nParallel != null ||" in page
|
||||
assert 'aria-label="Parallel decode slots"' in page
|
||||
api_types = _read("features/chat/types/api.ts")
|
||||
assert "n_parallel?: number | null;" in api_types
|
||||
runtime = _read("features/chat/hooks/use-chat-model-runtime.ts")
|
||||
# Click-time snapshot, /load body, validate preflight, cross-model reset and
|
||||
# failed-switch rollback all carry the value.
|
||||
assert "pendingLoadConfig?.nParallel" in runtime
|
||||
# GGUF-gated, like the compare pane: a transformers load has no slots.
|
||||
assert "n_parallel: isGguf ? loadNParallel : null," in runtime
|
||||
assert "n_parallel: validateNParallel," in runtime
|
||||
assert "loadNParallel = pendingLoadConfig?.nParallel ?? null;" in runtime
|
||||
assert "n_parallel: stateBeforeUnload.loadedNParallel," in runtime
|
||||
chat_api = _read("features/chat/api/chat-api.ts")
|
||||
assert "n_parallel: payload.n_parallel," in chat_api
|
||||
composer = _read("features/chat/shared-composer.tsx")
|
||||
# The compare pane is a second /load builder; its preflight sizes like its load.
|
||||
assert composer.count("n_parallel: ownConfig.nParallel ?? null,") == 2
|
||||
adapter = _read("features/chat/api/chat-adapter.ts")
|
||||
# The startup auto-load is a third builder reading the remembered config.
|
||||
assert adapter.count("n_parallel: config.nParallel ?? null,") == 2
|
||||
# ... and records it as loaded through the diffusion-gated local below.
|
||||
assert "loadedNParallel: committedSlots," in adapter
|
||||
status = _read("features/chat/lib/apply-inference-status-to-store.ts")
|
||||
# Hydration seeds the rollback BASELINE only; adopting the resolved echo into
|
||||
# the control would pin a blank "server default" to a number.
|
||||
assert "loadedNParallel: status.requested_parallel_slots," in status
|
||||
assert "nParallel: status.requested_parallel_slots," not in status
|
||||
sidebar = _read("features/model-picker/components/sidebar-model-config.tsx")
|
||||
# The sidebar form remounts when an external change lands.
|
||||
assert 'config.nParallel ?? "",' in sidebar
|
||||
|
||||
|
||||
def test_parallel_slots_control_cleared_when_the_load_never_sent_them():
|
||||
"""`nParallel` is the editable control ("blank = follow the server default")
|
||||
and `loadedNParallel` the rollback baseline. A success path that sends no
|
||||
slot count must blank the control, or a value staged for another model shows
|
||||
as applied, is persisted into this model's config (`isDefaultConfig` keys on
|
||||
nParallel) and is re-sent by the next Apply. Each assertion below is the only
|
||||
thing pinning one such path."""
|
||||
status = " ".join(_read("features/chat/lib/apply-inference-status-to-store.ts").split())
|
||||
# A model/variant swap underneath this tab must reset the control like
|
||||
# performLoad's cross-model reset, or model A's count follows onto model B.
|
||||
# Narrowly gated -- see test_hydration_keeps_the_slot_control_when_readopting_the_running_model.
|
||||
assert "...(seedLoadParams && slotsModelChanged && { nParallel: null })," in status
|
||||
# ... while still never adopting the RESOLVED echo into the control.
|
||||
assert "nParallel: status.requested_parallel_slots," not in status
|
||||
|
||||
adapter = _read("features/chat/api/chat-adapter.ts")
|
||||
# Slice the two success branches apart, bounding the second at the shared tail
|
||||
# so it cannot swallow the fresh-default path below and stay green.
|
||||
candidate = adapter.split("async function loadAutoLoadCandidate", 1)[1]
|
||||
gguf_branch, non_gguf_rest = candidate.split('if (candidate.kind === "gguf") {', 1)[1].split(
|
||||
"\n } else {\n", 1
|
||||
)
|
||||
non_gguf_branch = non_gguf_rest.split("if (!(loadResp.is_lora ?? false)) {", 1)[0]
|
||||
# The cached-GGUF branch keeps the remembered override via the gated local...
|
||||
assert "nParallel: committedSlots," in gguf_branch
|
||||
assert "nParallel: null," not in gguf_branch
|
||||
# ... the safetensors fallback sends no slots, so it clears both, or the count
|
||||
# survives on a model whose form does not even render the field.
|
||||
assert "nParallel: null," in non_gguf_branch
|
||||
assert "loadedNParallel: null," in non_gguf_branch
|
||||
|
||||
fresh_default = adapter.split("No downloaded models found. Fetching", 1)[1].split(
|
||||
'showAutoLoadSuccess("Loaded Qwen', 1
|
||||
)[0]
|
||||
# The fresh-default download omits the slots, so its success state clears both,
|
||||
# or the control reads as an unapplied edit against the seeded baseline.
|
||||
assert "n_parallel" not in fresh_default.split("saveSpeculativeType", 1)[0]
|
||||
assert "nParallel: null," in fresh_default
|
||||
assert "loadedNParallel: null," in fresh_default
|
||||
|
||||
|
||||
def test_hydration_clears_the_slot_baseline_for_a_slotless_model():
|
||||
"""The baseline is what a rollback re-sends and what preset capture reads, so
|
||||
a model that cannot have slots must not inherit the previous GGUF's count.
|
||||
/status omits the echo for non-GGUF and sends an explicit null for diffusion;
|
||||
an absent field on a GGUF is an older backend and must NOT wipe it."""
|
||||
src = _read("features/chat/lib/apply-inference-status-to-store.ts")
|
||||
assert (
|
||||
"(status.is_gguf === false || status.requested_parallel_slots === null) && {" in src
|
||||
), "the slotless clear must key on is_gguf or an explicit null echo"
|
||||
clear = src.index("status.is_gguf === false || status.requested_parallel_slots === null")
|
||||
assert "loadedNParallel: null," in src[clear : clear + 200]
|
||||
# Never `!= null`: that also matches the absent field an older backend sends.
|
||||
assert "status.requested_parallel_slots !== null && {" not in src
|
||||
|
||||
|
||||
def test_hydration_keeps_the_slot_control_when_readopting_the_running_model():
|
||||
"""`hydratingExistingModel` is true whenever the incoming status disagrees
|
||||
with what this tab last recorded, which includes RE-ADOPTING a model the tab
|
||||
never lost: the resident-adopt branch restores the model's own per-model
|
||||
config and only then hydrates, passing the EXTERNAL id as
|
||||
`previousCheckpoint`. An ungated clear there wipes the slot count that branch
|
||||
just restored, and the blank persists into `savePerModelConfig`, so a Save
|
||||
the user reads as a no-op erases their remembered override.
|
||||
|
||||
Only that branch knows the model is unchanged, so it says so explicitly.
|
||||
Slot counts cannot stand in: the echo falls back to the server-wide default,
|
||||
so a genuine A->B swap can echo exactly A's explicit count."""
|
||||
status = " ".join(_read("features/chat/lib/apply-inference-status-to-store.ts").split())
|
||||
assert (
|
||||
"const slotsModelChanged = hydratingExistingModel && !options.readoptingSameModel;"
|
||||
in status
|
||||
)
|
||||
assert "...(seedLoadParams && slotsModelChanged && { nParallel: null })," in status
|
||||
# Never a slot-count proxy for "same model".
|
||||
assert "prevState.loadedNParallel === (status.requested_parallel_slots" not in status
|
||||
# The baseline seed stays ungated, or a rollback after a tab reload restores
|
||||
# the model at the server default slots.
|
||||
assert "loadedNParallel: status.requested_parallel_slots," in status
|
||||
|
||||
runtime = " ".join(_read("features/chat/hooks/use-chat-model-runtime.ts").split())
|
||||
resident = runtime.split("if (!forceReload && isExternalModelId(selectedCheckpoint)) {", 1)[
|
||||
1
|
||||
].split("const stopDecision", 1)[0]
|
||||
# What makes the scenario reachable: the branch restores the model's own
|
||||
# config, then hydrates against the external id.
|
||||
assert "applyPerModelConfigToRuntime(selection.previousConfig);" in resident
|
||||
assert "previousCheckpoint: selectedCheckpoint," in resident
|
||||
# Only reachable because the branch matched the id AND the variant first.
|
||||
assert "resolveInferenceCheckpointId(residentStatus) === modelId" in resident
|
||||
assert "readoptingSameModel: true," in resident
|
||||
# The refresh() hydrate must NOT claim it: there the model really can change.
|
||||
poll = runtime.split("setModels(listRes.models.map(toChatModelSummary));", 1)[1].split(
|
||||
"} else if (!statusRes.active_model", 1
|
||||
)[0]
|
||||
assert "applyActiveModelStatusToStore(statusRes, {" in poll
|
||||
assert "readoptingSameModel" not in poll
|
||||
|
||||
|
||||
def test_parallel_slots_are_never_recorded_for_a_diffusion_load():
|
||||
"""A DiffusionGemma GGUF answers ``is_gguf: true``, but its runner ignores
|
||||
``--parallel``, so ``_parallel_slot_echo`` reports null slots for it. The
|
||||
three load success paths must gate on ``is_diffusion`` too, or they record a
|
||||
click-time count the load never committed.
|
||||
|
||||
That phantom does not stay put: ``capturePresetLoadConfig`` snapshots
|
||||
``nParallel`` with no model gate and a preset carries no model identity, so
|
||||
applying it over a TEXT GGUF sends the count as a real ``n_parallel``.
|
||||
"""
|
||||
runtime = " ".join(_read("features/chat/hooks/use-chat-model-runtime.ts").split())
|
||||
# One gated local feeds the control and the baseline, so they cannot drift.
|
||||
assert "(loadResponse.is_gguf ?? false) && !(loadResponse.is_diffusion ?? false)" in runtime
|
||||
assert "nParallel: committedSlots," in runtime
|
||||
assert "loadedNParallel: committedSlots," in runtime
|
||||
|
||||
adapter = " ".join(_read("features/chat/api/chat-adapter.ts").split())
|
||||
assert (
|
||||
"const committedSlots = (loadResp.is_diffusion ?? false) ? null "
|
||||
": (config.nParallel ?? null);" in adapter
|
||||
)
|
||||
assert "nParallel: committedSlots," in adapter
|
||||
assert "loadedNParallel: committedSlots," in adapter
|
||||
|
||||
composer = " ".join(_read("features/chat/shared-composer.tsx").split())
|
||||
assert "targetIsGguf && !(resp.is_diffusion ?? false)" in composer
|
||||
assert "nParallel: committedSlots," in composer
|
||||
assert "loadedNParallel: committedSlots," in composer
|
||||
|
||||
|
||||
def test_hydration_restores_a_remembered_slot_override():
|
||||
"""The control is never seeded from the status echo, so a model running on a
|
||||
remembered override shows a BLANK slot control after a browser reload or a
|
||||
tab move to another GGUF. `ModelConfigPage.resolveInitial` prefers the live
|
||||
store for the active model, so that blank is what the form edits: the next
|
||||
Apply reloads at the server default and a Save writes the blank over the
|
||||
remembered count.
|
||||
|
||||
The seed is deliberately narrow: storage is read only on a fresh store or a
|
||||
model change, never on a steady poll, and the value is adopted only when the
|
||||
server already runs that exact count, which proves it is this model's own.
|
||||
"""
|
||||
src = _read("features/chat/lib/apply-inference-status-to-store.ts")
|
||||
status = " ".join(src.split())
|
||||
assert (
|
||||
"resolveInitialConfig(checkpointId, status.gguf_variant ?? null)" in status
|
||||
), "the remembered override comes from per-model storage, not the echo"
|
||||
assert (
|
||||
"const slotsUnseeded = prevState.loadedNParallel === null && "
|
||||
"prevState.nParallel === null;" in status
|
||||
)
|
||||
assert (
|
||||
"status.is_gguf && (slotsUnseeded || slotsModelChanged)" in status
|
||||
), "storage is read on a fresh store or a model change, never on a steady poll"
|
||||
assert (
|
||||
"...(seedLoadParams && (slotsUnseeded || slotsModelChanged) &&" in status
|
||||
), "the seed fires in both cases the clear leaves the control blank"
|
||||
assert (
|
||||
"rememberedNParallel != null && rememberedNParallel === "
|
||||
"status.requested_parallel_slots && { nParallel: rememberedNParallel, }" in status
|
||||
)
|
||||
# Both cases trip the model-change clear, so the seed only survives by
|
||||
# being spread after it.
|
||||
assert src.index("slotsModelChanged && { nParallel: null }") < src.index(
|
||||
"nParallel: rememberedNParallel,"
|
||||
)
|
||||
|
||||
|
||||
def test_failed_switch_rollback_restores_the_slot_intent_not_the_resolved_count():
|
||||
"""`loadedNParallel` holds a RESOLVED count even for a load that sent no
|
||||
slots (the echo falls back to the server-wide default), so it is the right
|
||||
value to re-send when recreating the previous server and the wrong one to put
|
||||
back in the control: it turns "follow the server default" into an explicit
|
||||
override that a later Save or preset capture pins. The outer catch only
|
||||
repairs that for a staged config, so a plain string pick keeps the phantom.
|
||||
|
||||
The intent comes from the picker's own pre-switch snapshot when there is one:
|
||||
chat-page pre-applies the TARGET's config before calling selectModel, so the
|
||||
live control describes the outgoing model only for a bare pick."""
|
||||
runtime = " ".join(_read("features/chat/hooks/use-chat-model-runtime.ts").split())
|
||||
assert (
|
||||
'const previousNParallel = typeof selection !== "string" && '
|
||||
"selection.previousConfig ? (selection.previousConfig.nParallel ?? null) "
|
||||
": useChatRuntimeStore.getState().nParallel;" in runtime
|
||||
)
|
||||
assert runtime.index("const previousNParallel") < runtime.index(
|
||||
"applyPerModelConfigToRuntime(pendingLoadConfig);"
|
||||
), "a config staged on the selection must not replace it either"
|
||||
picker = " ".join(_read("features/chat/chat-page.tsx").split())
|
||||
assert (
|
||||
"const previousConfig = currentRuntimePerModelConfig({ includeMaxSeqLength: true, }); "
|
||||
"const hasAppliedConfig = applyModelLoadConfigToRuntime(" in picker
|
||||
), "the snapshot must be taken before the target's config is applied"
|
||||
rollback = runtime.split("const rollbackSpeculativeType", 1)[1]
|
||||
assert "nParallel: previousNParallel," in rollback
|
||||
# Baseline and reload payload keep the resolved count, or the rollback
|
||||
# recreates the previous model at a different slot count.
|
||||
assert "loadedNParallel: stateBeforeUnload.loadedNParallel ?? null," in rollback
|
||||
assert "n_parallel: stateBeforeUnload.loadedNParallel," in runtime
|
||||
|
||||
|
||||
def test_vulkan_inference_devices_are_the_pickable_set():
|
||||
"""GGUF loads run through llama-server, so on a Vulkan build the picker must
|
||||
offer the inference inventory (ggml ordinals, the space `--device Vulkan<i>`
|
||||
|
|
|
|||
188
tests/studio/test_tauri_deep_link_contract.py
Normal file
188
tests/studio/test_tauri_deep_link_contract.py
Normal file
|
|
@ -0,0 +1,188 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""Desktop deep-link configuration, routing, and validation contracts."""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
import shutil
|
||||
import subprocess
|
||||
import textwrap
|
||||
|
||||
import pytest
|
||||
|
||||
try:
|
||||
import tomllib
|
||||
except ModuleNotFoundError:
|
||||
tomllib = pytest.importorskip("tomli")
|
||||
|
||||
|
||||
REPO = Path(__file__).resolve().parents[2]
|
||||
FRONTEND = REPO / "studio/frontend"
|
||||
TAURI = REPO / "studio/src-tauri"
|
||||
PARSER = FRONTEND / "src/features/deep-links/parse-deep-link.ts"
|
||||
INTENT_GATE = FRONTEND / "src/features/deep-links/deep-link-intent.ts"
|
||||
GGUF_FILENAME = FRONTEND / "src/features/hub/lib/gguf-filename.ts"
|
||||
|
||||
|
||||
def test_unsloth_deep_link_parser_guardrails(tmp_path: Path) -> None:
|
||||
if shutil.which("node") is None:
|
||||
pytest.skip("node not available")
|
||||
probe = subprocess.run(
|
||||
["node", "--experimental-strip-types", "--version"],
|
||||
capture_output = True,
|
||||
text = True,
|
||||
timeout = 5,
|
||||
)
|
||||
if probe.returncode != 0:
|
||||
pytest.skip("node --experimental-strip-types not available")
|
||||
|
||||
(tmp_path / "parse-deep-link.ts").write_text(
|
||||
PARSER.read_text(encoding = "utf-8"), encoding = "utf-8"
|
||||
)
|
||||
|
||||
(tmp_path / "gguf-filename.ts").write_text(
|
||||
GGUF_FILENAME.read_text(encoding = "utf-8"), encoding = "utf-8"
|
||||
)
|
||||
|
||||
(tmp_path / "deep-link-intent.ts").write_text(
|
||||
INTENT_GATE.read_text(encoding = "utf-8"), encoding = "utf-8"
|
||||
)
|
||||
script = textwrap.dedent("""
|
||||
import assert from "node:assert/strict";
|
||||
import { parseUnslothDeepLink } from "./parse-deep-link.ts";
|
||||
|
||||
import { createDeepLinkIntentGate } from "./deep-link-intent.ts";
|
||||
import {
|
||||
ggufFilenamesMatch,
|
||||
ggufSelectionOverrideMatchesIntent,
|
||||
} from "./gguf-filename.ts";
|
||||
|
||||
const valid = new Map([
|
||||
[
|
||||
"unsloth://open_from_hf?model=unsloth/Laguna-S-2.1-GGUF",
|
||||
{ model: "unsloth/Laguna-S-2.1-GGUF" },
|
||||
],
|
||||
[
|
||||
"unsloth://open_from_hf/?model=org/repo_name",
|
||||
{ model: "org/repo_name" },
|
||||
],
|
||||
[
|
||||
"unsloth://open_from_hf?model=org%2Frepo",
|
||||
{ model: "org/repo" },
|
||||
],
|
||||
[
|
||||
"unsloth://open_from_hf?model=unsloth/Laguna-S-2.1-GGUF&file=Laguna-S-2.1-UD-IQ3_XXS.gguf",
|
||||
{
|
||||
model: "unsloth/Laguna-S-2.1-GGUF",
|
||||
file: "Laguna-S-2.1-UD-IQ3_XXS.gguf",
|
||||
},
|
||||
],
|
||||
[
|
||||
"unsloth://open_from_hf?file=weights%2Fmodel-Q4_K_M.gguf&model=org/repo",
|
||||
{ model: "org/repo", file: "weights/model-Q4_K_M.gguf" },
|
||||
],
|
||||
[
|
||||
`unsloth://open_from_hf?model=${"a".repeat(96)}/${"b".repeat(96)}`,
|
||||
{ model: `${"a".repeat(96)}/${"b".repeat(96)}` },
|
||||
],
|
||||
]);
|
||||
for (const [url, intent] of valid) {
|
||||
assert.deepEqual(parseUnslothDeepLink(url), intent, url);
|
||||
}
|
||||
|
||||
assert.equal(
|
||||
ggufFilenamesMatch(
|
||||
"weights/model-Q4_K_M-00002-of-00002.gguf",
|
||||
"weights/model-Q4_K_M-00001-of-00002.gguf",
|
||||
),
|
||||
true,
|
||||
);
|
||||
assert.equal(
|
||||
ggufFilenamesMatch("model-Q4_K_M.GGUF", "model-q4_k_m.gguf"),
|
||||
true,
|
||||
);
|
||||
assert.equal(ggufFilenamesMatch("mmproj-F16.gguf", "model-F16.gguf"), false);
|
||||
|
||||
assert.equal(ggufSelectionOverrideMatchesIntent("a.gguf", 2, "a.gguf", 2), true);
|
||||
assert.equal(ggufSelectionOverrideMatchesIntent("a.gguf", 2, "a.gguf", 1), false);
|
||||
|
||||
let now = 1_000;
|
||||
const acceptIntent = createDeepLinkIntentGate(2_000, () => now);
|
||||
assert.equal(acceptIntent("org/repo", "a.gguf"), 1);
|
||||
assert.equal(acceptIntent("org/repo", "a.gguf"), null);
|
||||
assert.equal(acceptIntent("org/repo", "b.gguf"), 2);
|
||||
now = 3_000;
|
||||
assert.equal(acceptIntent("org/repo", "b.gguf"), 3);
|
||||
|
||||
|
||||
const invalid = [
|
||||
"",
|
||||
"https://open_from_hf?model=org/repo",
|
||||
"UNSLOTH://open_from_hf?model=org/repo",
|
||||
"unsloth://OPEN_FROM_HF?model=org/repo",
|
||||
"unsloth://open_from_hf/path?model=org/repo",
|
||||
"unsloth://open_from_hf/%2e%2e?model=org/repo",
|
||||
"unsloth://user@open_from_hf?model=org/repo",
|
||||
"unsloth://open_from_hf:42?model=org/repo",
|
||||
"unsloth://open_from_hf?model=org/repo#fragment",
|
||||
"unsloth://open_from_hf?model=org/repo&download=true",
|
||||
|
||||
"unsloth://open_from_hf?model=org/repo&file=model.gguf&file=other.gguf",
|
||||
"unsloth://open_from_hf?model=org/repo&file=",
|
||||
"unsloth://open_from_hf?model=org/repo&file=../model.gguf",
|
||||
"unsloth://open_from_hf?model=org/repo&file=%2Fmodel.gguf",
|
||||
"unsloth://open_from_hf?model=org/repo&file=model.safetensors",
|
||||
"unsloth://open_from_hf?model=org/repo&model=other/repo",
|
||||
"unsloth://open_from_hf?model=repo",
|
||||
"unsloth://open_from_hf?model=org/repo/extra",
|
||||
"unsloth://open_from_hf?model=-org/repo",
|
||||
"unsloth://open_from_hf?model=org/repo.",
|
||||
|
||||
"unsloth://open_from_hf?model=org/repo.git",
|
||||
"unsloth://open_from_hf?model=org/repo--name",
|
||||
"unsloth://open_from_hf?model=org/repo..name",
|
||||
];
|
||||
for (const url of invalid) {
|
||||
assert.equal(parseUnslothDeepLink(url), null, url);
|
||||
}
|
||||
""")
|
||||
result = subprocess.run(
|
||||
["node", "--experimental-strip-types", "--no-warnings", "--input-type=module"],
|
||||
input = script,
|
||||
cwd = tmp_path,
|
||||
capture_output = True,
|
||||
text = True,
|
||||
timeout = 30,
|
||||
)
|
||||
assert result.returncode == 0, f"stderr: {result.stderr}\nstdout: {result.stdout}"
|
||||
|
||||
|
||||
def test_tauri_registers_only_the_unsloth_scheme() -> None:
|
||||
cargo = tomllib.loads((TAURI / "Cargo.toml").read_text(encoding = "utf-8"))
|
||||
dependencies = cargo["dependencies"]
|
||||
assert "tauri-plugin-deep-link" in dependencies
|
||||
single_instance = dependencies["tauri-plugin-single-instance"]
|
||||
assert isinstance(single_instance, dict)
|
||||
assert "deep-link" in single_instance.get("features", [])
|
||||
|
||||
config = json.loads((TAURI / "tauri.conf.json").read_text(encoding = "utf-8"))
|
||||
assert config["plugins"]["deep-link"]["desktop"]["schemes"] == ["unsloth"]
|
||||
|
||||
capabilities = json.loads((TAURI / "capabilities/default.json").read_text(encoding = "utf-8"))
|
||||
assert "deep-link:default" in capabilities["permissions"]
|
||||
assert "core:window:allow-unminimize" in capabilities["permissions"]
|
||||
|
||||
main = (TAURI / "src/main.rs").read_text(encoding = "utf-8")
|
||||
assert main.index("tauri_plugin_single_instance::init") < main.index(
|
||||
"tauri_plugin_deep_link::init()"
|
||||
)
|
||||
assert "DeepLinkExt" in main
|
||||
assert "if let Err(error) = app.deep_link().register_all()" in main
|
||||
assert 'warn!("Failed to register deep-link handlers: {error}")' in main
|
||||
assert 'target_os = "linux"' in main
|
||||
desktop_template = TAURI / "linux/unsloth.desktop"
|
||||
assert config["bundle"]["linux"]["deb"]["desktopTemplate"] == "./linux/unsloth.desktop"
|
||||
desktop = desktop_template.read_text(encoding = "utf-8")
|
||||
assert "Exec={{exec}} %u" in desktop
|
||||
assert "MimeType=x-scheme-handler/unsloth;" in desktop
|
||||
|
|
@ -11,7 +11,11 @@ dequant reference.
|
|||
import pytest
|
||||
import torch
|
||||
|
||||
pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason = "needs CUDA")
|
||||
cuda_available = torch.cuda.is_available()
|
||||
xpu_available = hasattr(torch, "xpu") and torch.xpu.is_available()
|
||||
dev = "cuda" if cuda_available else "xpu" if xpu_available else "cpu"
|
||||
|
||||
pytestmark = pytest.mark.skipif(not (cuda_available or xpu_available), reason = "needs CUDA or XPU")
|
||||
|
||||
|
||||
def _reference(X, weight, scale, block):
|
||||
|
|
@ -27,7 +31,6 @@ def test_tiny_non_tileable_forward_backward_matches_reference():
|
|||
from unsloth.kernels.fp8 import FP8BlockQuantLinear
|
||||
|
||||
torch.manual_seed(0)
|
||||
dev = "cuda"
|
||||
block = [128, 128]
|
||||
m, n = 8, 8 # non-tileable, in-dim % 128 != 0
|
||||
weight = torch.randn(m, n, device = dev, dtype = torch.bfloat16) # (out=m, in=n)
|
||||
|
|
@ -50,7 +53,6 @@ def test_e8m0_scale_is_upcast_and_runs():
|
|||
if not hasattr(torch, "float8_e8m0fnu"):
|
||||
pytest.skip("torch build lacks float8_e8m0fnu")
|
||||
|
||||
dev = "cuda"
|
||||
m, n = 8, 8
|
||||
weight = torch.randn(m, n, device = dev, dtype = torch.bfloat16)
|
||||
scale = (torch.rand(1, 1, device = dev) + 1.0).to(torch.float8_e8m0fnu)
|
||||
|
|
@ -70,7 +72,6 @@ def test_rectangular_block_dequant_matches_reference():
|
|||
from unsloth.kernels.fp8 import _blockwise_weight_dequant_any_shape
|
||||
|
||||
torch.manual_seed(0)
|
||||
dev = "cuda"
|
||||
block = [64, 128]
|
||||
m, n = 64, 256 # evenly tiled: 64 % 64 == 0, 256 % 128 == 0
|
||||
weight = torch.randn(m, n, device = dev, dtype = torch.bfloat16)
|
||||
|
|
@ -94,7 +95,6 @@ def test_e8m0_scale_preserves_non_default_block_size_attr():
|
|||
pytest.skip("torch build lacks float8_e8m0fnu")
|
||||
|
||||
torch.manual_seed(0)
|
||||
dev = "cuda"
|
||||
block = [64, 64]
|
||||
# in-dim 96 is not divisible by block[1]=64 -> forward takes the torch dequant
|
||||
# fallback (no fp8 matmul kernel). Scale shape (2, 2) validates for [64, 64] but
|
||||
|
|
|
|||
|
|
@ -14,6 +14,11 @@ Two gaps it misses:
|
|||
future rename): reading them with ``[]`` raises ``KeyError`` into the bare ``except``,
|
||||
taking the 4bit half, the probe's whole purpose, down with it.
|
||||
|
||||
Both of the above only reach the block table. The last two tests take the row branch, which
|
||||
``load_in_fp8 = True`` plus ``UNSLOTH_HAS_FBGEMM`` selects ahead of block: deleting that branch,
|
||||
or dropping ``_resolve_with_mappers``' ``fp8_row`` argument so it falls back to the installed
|
||||
table, both leave every other test here green.
|
||||
|
||||
``loader_utils`` imports torch, so ast-extract the resolvers and run them against a stubbed
|
||||
``requests``, as in ``tests/test_bad_mappings_redirect.py``.
|
||||
"""
|
||||
|
|
@ -33,6 +38,8 @@ _NEW_OFFICIAL = "zeta-org/Zeta-9B-Only-On-Main-FP8"
|
|||
_NEW_BLOCK = "unsloth/Zeta-9B-Only-On-Main-FP8-Block"
|
||||
_NEW_ROW = "unsloth/Zeta-9B-Only-On-Main-FP8-Row"
|
||||
_ANCHOR = ' "unsloth/Kimi-K2-Instruct-BF16" : ('
|
||||
# Row table only, so the block branch cannot answer for it and mask a row-path regression.
|
||||
_ROW_ONLY = "zeta-org/Zeta-9B-Row-Only-FP8"
|
||||
|
||||
|
||||
def _mapper_source():
|
||||
|
|
@ -51,6 +58,11 @@ def _with_extra_fp8_model(source):
|
|||
return source.replace(_ANCHOR, entry + _ANCHOR, 1)
|
||||
|
||||
|
||||
def _with_row_only_fp8_model(source):
|
||||
"""Fetched row table only. Block must not know it, or the block branch answers instead."""
|
||||
return source + f'\nFLOAT_TO_FP8_ROW_MAPPER["{_ROW_ONLY.lower()}"] = "{_NEW_ROW}"\n'
|
||||
|
||||
|
||||
def _without_fp8_tables(source):
|
||||
"""A mapper.py from before the fp8 tables existed."""
|
||||
return source.replace("FLOAT_TO_FP8_BLOCK_MAPPER", "SOME_OTHER_BLOCK_TABLE").replace(
|
||||
|
|
@ -153,3 +165,44 @@ def test_probe_survives_a_fetched_mapper_without_the_fp8_tables(monkeypatch):
|
|||
assert (
|
||||
int_to_float and float_to_int and map_to_16bit
|
||||
), "a fetched mapper.py without the fp8 tables must not take the 4bit upgrade check down"
|
||||
|
||||
|
||||
def test_fbgemm_prefers_the_row_table_over_the_block_one(monkeypatch):
|
||||
"""With FBGEMM, `load_in_fp8 = True` must resolve row-scaled, not blockwise."""
|
||||
monkeypatch.setenv("UNSLOTH_HAS_FBGEMM", "1")
|
||||
namespace = _load_resolver(_mapper_source())
|
||||
row = namespace["FLOAT_TO_FP8_ROW_MAPPER"]
|
||||
block = namespace["FLOAT_TO_FP8_BLOCK_MAPPER"]
|
||||
|
||||
key = next(k for k in row if k in block and row[k] != block[k])
|
||||
resolved = namespace["get_model_name"](key, load_in_4bit = False, load_in_fp8 = True)
|
||||
|
||||
assert resolved == row[key], (
|
||||
f"FBGEMM must take the row branch for {key!r}, got {resolved!r} "
|
||||
f"(the blockwise answer is {block[key]!r})"
|
||||
)
|
||||
|
||||
|
||||
def test_probe_answers_for_a_row_only_repo_the_fetched_mapper_knows(monkeypatch):
|
||||
"""The row half of the probe needs the FETCHED row table, same as the block half."""
|
||||
monkeypatch.setenv("UNSLOTH_HAS_FBGEMM", "1")
|
||||
installed = _mapper_source()
|
||||
namespace = _load_resolver(installed)
|
||||
installed_row = namespace["FLOAT_TO_FP8_ROW_MAPPER"]
|
||||
key = _ROW_ONLY.lower()
|
||||
assert key not in installed_row, "the installed row table must not know it"
|
||||
assert key not in namespace["FLOAT_TO_FP8_BLOCK_MAPPER"], "no block entry, or block answers"
|
||||
|
||||
_install_fake_requests(monkeypatch, _with_row_only_fp8_model(installed))
|
||||
_install_fake_vllm_absent(monkeypatch, namespace)
|
||||
|
||||
try:
|
||||
resolved = namespace["get_model_name"](_ROW_ONLY, load_in_4bit = False, load_in_fp8 = True)
|
||||
except NotImplementedError as error:
|
||||
assert "not supported in your current Unsloth version" in str(error)
|
||||
else:
|
||||
raise AssertionError(
|
||||
f"a fetched-only row-scaled repo must raise the upgrade error, got {resolved!r}"
|
||||
)
|
||||
|
||||
assert namespace["FLOAT_TO_FP8_ROW_MAPPER"] is installed_row
|
||||
|
|
|
|||
128
tests/test_raw_text_json_loading.py
Normal file
128
tests/test_raw_text_json_loading.py
Normal file
|
|
@ -0,0 +1,128 @@
|
|||
"""Regression test for .json parsing in unsloth/dataprep/raw_text.py.
|
||||
|
||||
Both .json and .jsonl map to the "json_lines" handler, which used to parse the
|
||||
file one line at a time. A real .json file is a single JSON document (commonly
|
||||
a top-level list of records), so every line failed json.loads, the whole
|
||||
document was dropped, and the handler returned "" (load_from_file then rejected
|
||||
the valid file as "empty"). The handler now parses the file as one JSON value
|
||||
first and falls back to line-by-line for true .jsonl.
|
||||
|
||||
raw_text.py's only third-party import is `datasets`, so we stub it and exec the
|
||||
module directly, with no `import unsloth` (which needs a GPU / unsloth_zoo).
|
||||
"""
|
||||
|
||||
import json
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
|
||||
RAW_TEXT_PATH = Path(__file__).parents[1] / "unsloth" / "dataprep" / "raw_text.py"
|
||||
|
||||
|
||||
def _load_raw_text():
|
||||
sys.modules.setdefault("datasets", types.SimpleNamespace(Dataset = object))
|
||||
module = types.ModuleType("unsloth_raw_text_under_test")
|
||||
exec(
|
||||
compile(RAW_TEXT_PATH.read_text(encoding = "utf-8"), str(RAW_TEXT_PATH), "exec"),
|
||||
module.__dict__,
|
||||
)
|
||||
return module
|
||||
|
||||
|
||||
def test_json_document_is_parsed_whole(tmp_path):
|
||||
loader = _load_raw_text().RawTextDataLoader(tokenizer = object())
|
||||
path = tmp_path / "data.json"
|
||||
path.write_text(
|
||||
json.dumps([{"text": "hello world"}, {"text": "second sample"}], indent = 2), encoding = "utf-8"
|
||||
)
|
||||
assert loader._read_file_by_format(str(path), "json_lines") == "hello world\n\nsecond sample"
|
||||
|
||||
|
||||
def test_jsonl_is_still_parsed_line_by_line(tmp_path):
|
||||
loader = _load_raw_text().RawTextDataLoader(tokenizer = object())
|
||||
path = tmp_path / "data.jsonl"
|
||||
path.write_text('{"text": "a"}\n{"text": "b"}\n', encoding = "utf-8")
|
||||
assert loader._read_file_by_format(str(path), "json_lines") == "a\n\nb"
|
||||
|
||||
|
||||
def test_jsonl_is_never_materialized(tmp_path):
|
||||
"""A .jsonl file must keep streaming, whole-document parsing is only for .json."""
|
||||
real_open = open
|
||||
|
||||
class _StreamOnlyFile:
|
||||
"""File wrapper that fails the test if the whole file is pulled into memory."""
|
||||
|
||||
def __init__(self, handle):
|
||||
self.handle = handle
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc_info):
|
||||
self.handle.close()
|
||||
return False
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self.handle)
|
||||
|
||||
def read(self, *args, **kwargs):
|
||||
raise AssertionError(".jsonl was read whole instead of streamed line by line")
|
||||
|
||||
def seek(self, *args, **kwargs):
|
||||
raise AssertionError(".jsonl was re-read instead of streamed line by line")
|
||||
|
||||
module = _load_raw_text()
|
||||
module.open = lambda *args, **kwargs: _StreamOnlyFile(real_open(*args, **kwargs))
|
||||
|
||||
path = tmp_path / "big.jsonl"
|
||||
path.write_text('{"text": "a"}\n\n{"text": "b"}\nnot json at all\n', encoding = "utf-8")
|
||||
loader = module.RawTextDataLoader(tokenizer = object())
|
||||
assert loader._read_file_by_format(str(path), "json_lines") == "a\n\nb"
|
||||
|
||||
|
||||
def test_json_holding_json_lines_still_falls_back(tmp_path):
|
||||
"""A .json file that actually holds JSON Lines still parses, via the per-line fallback."""
|
||||
loader = _load_raw_text().RawTextDataLoader(tokenizer = object())
|
||||
path = tmp_path / "mislabelled.json"
|
||||
path.write_text('{"text": "a"}\n{"text": "b"}\n', encoding = "utf-8")
|
||||
assert loader._read_file_by_format(str(path), "json_lines") == "a\n\nb"
|
||||
|
||||
|
||||
def test_utf8_bom_json_document_is_parsed(tmp_path):
|
||||
"""Windows tooling prefixes a UTF-8 BOM; it must not sink the whole document."""
|
||||
loader = _load_raw_text().RawTextDataLoader(tokenizer = object())
|
||||
path = tmp_path / "bom.json"
|
||||
path.write_text(
|
||||
json.dumps([{"text": "hello world"}, {"text": "second sample"}], indent = 2),
|
||||
encoding = "utf-8-sig",
|
||||
)
|
||||
assert path.read_bytes().startswith(b"\xef\xbb\xbf")
|
||||
assert loader._read_file_by_format(str(path), "json_lines") == "hello world\n\nsecond sample"
|
||||
|
||||
|
||||
def test_utf8_bom_jsonl_keeps_first_record(tmp_path):
|
||||
"""A BOM must not silently drop the first .jsonl record."""
|
||||
loader = _load_raw_text().RawTextDataLoader(tokenizer = object())
|
||||
path = tmp_path / "bom.jsonl"
|
||||
path.write_text('{"text": "a"}\n{"text": "b"}\n', encoding = "utf-8-sig")
|
||||
assert loader._read_file_by_format(str(path), "json_lines") == "a\n\nb"
|
||||
|
||||
|
||||
def test_utf8_bom_json_holding_json_lines_falls_back(tmp_path):
|
||||
"""The per-line fallback re-reads from byte 0, so the BOM must be stripped again."""
|
||||
loader = _load_raw_text().RawTextDataLoader(tokenizer = object())
|
||||
path = tmp_path / "bom_mislabelled.json"
|
||||
path.write_text('{"text": "a"}\n{"text": "b"}\n', encoding = "utf-8-sig")
|
||||
assert loader._read_file_by_format(str(path), "json_lines") == "a\n\nb"
|
||||
|
||||
|
||||
def test_utf8_bom_plain_text_and_csv(tmp_path):
|
||||
"""The BOM also leaks into .txt training text and the first .csv column name."""
|
||||
loader = _load_raw_text().RawTextDataLoader(tokenizer = object())
|
||||
txt = tmp_path / "bom.txt"
|
||||
txt.write_text("hello", encoding = "utf-8-sig")
|
||||
assert loader._read_file_by_format(str(txt), "plain_text") == "hello"
|
||||
|
||||
csv_path = tmp_path / "bom.csv"
|
||||
csv_path.write_text("text,other\nhello,x\n", encoding = "utf-8-sig")
|
||||
assert loader._read_file_by_format(str(csv_path), "csv_text_column") == "hello"
|
||||
|
|
@ -2,6 +2,9 @@ from tqdm import tqdm
|
|||
import torch
|
||||
import pandas as pd
|
||||
|
||||
# DEVICE_TYPE_TORCH, not DEVICE_TYPE: the latter can be "hip"/"mlx", which .to() rejects.
|
||||
from unsloth.device_type import DEVICE_TYPE_TORCH
|
||||
|
||||
model_comparison_results = {}
|
||||
|
||||
|
||||
|
|
@ -17,7 +20,7 @@ def ppl_model(model, tokenizer, dataset):
|
|||
for begin_loc in range(0, seq_len, stride):
|
||||
end_loc = min(begin_loc + max_length, seq_len)
|
||||
trg_len = end_loc - prev_end_loc
|
||||
input_ids = encodings.input_ids[:, begin_loc:end_loc].to("cuda")
|
||||
input_ids = encodings.input_ids[:, begin_loc:end_loc].to(DEVICE_TYPE_TORCH)
|
||||
target_ids = input_ids.clone()
|
||||
target_ids[:, :-trg_len] = -100
|
||||
pad_token_id = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else 0
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Greedy generation in a left-padded batch must match solo batch-size-1
|
|||
generation for the first PREFIX_TOKENS tokens (the bug makes padded rows
|
||||
diverge into garbage immediately; a full-length match would be flaky due to
|
||||
benign batch-numerics tie-flips deep in the sequence) and must not be
|
||||
gibberish. Skipped without CUDA. Run: `python -m pytest
|
||||
gibberish. Skipped without a GPU. Run: `python -m pytest
|
||||
tests/utils/test_batched_leftpad_generation_gpu.py -v`.
|
||||
"""
|
||||
|
||||
|
|
@ -12,8 +12,19 @@ import pytest
|
|||
import torch
|
||||
|
||||
cuda_available = torch.cuda.is_available()
|
||||
xpu_available = hasattr(torch, "xpu") and torch.xpu.is_available()
|
||||
device = "cuda" if cuda_available else "xpu" if xpu_available else "cpu"
|
||||
|
||||
pytestmark = pytest.mark.skipif(not cuda_available, reason = "requires a CUDA GPU")
|
||||
# Non-strict rather than CUDA-only: keeps the XPU divergence visible, and goes
|
||||
# green by itself once XPU generation is fixed.
|
||||
pytestmark = [
|
||||
pytest.mark.skipif(not (cuda_available or xpu_available), reason = "requires a CUDA or XPU GPU"),
|
||||
pytest.mark.xfail(
|
||||
xpu_available and not cuda_available,
|
||||
reason = "batched left-padded generation diverges on XPU",
|
||||
strict = False,
|
||||
),
|
||||
]
|
||||
|
||||
MODEL_NAME = "unsloth/Qwen2.5-0.5B-Instruct"
|
||||
MAX_NEW_TOKENS = 32
|
||||
|
|
@ -53,7 +64,7 @@ def _chat(tokenizer, prompt):
|
|||
|
||||
def _generate(model, tokenizer, texts):
|
||||
inputs = tokenizer(texts, return_tensors = "pt", padding = True, add_special_tokens = False).to(
|
||||
"cuda"
|
||||
device
|
||||
)
|
||||
with torch.inference_mode():
|
||||
out = model.generate(
|
||||
|
|
|
|||
|
|
@ -44,6 +44,8 @@ def _build_packed_training_setup(tmp_path, device):
|
|||
dtype = torch.bfloat16
|
||||
else:
|
||||
dtype = torch.float16
|
||||
elif device.type == "xpu":
|
||||
dtype = torch.bfloat16
|
||||
|
||||
try:
|
||||
model, tokenizer = FastLanguageModel.from_pretrained(
|
||||
|
|
@ -76,8 +78,8 @@ def _build_packed_training_setup(tmp_path, device):
|
|||
max_length = 64,
|
||||
logging_steps = 1,
|
||||
max_steps = 1,
|
||||
fp16 = device.type == "cuda" and not torch.cuda.is_bf16_supported(),
|
||||
bf16 = device.type == "cuda" and torch.cuda.is_bf16_supported(),
|
||||
fp16 = dtype == torch.float16,
|
||||
bf16 = dtype == torch.bfloat16,
|
||||
dataset_num_proc = 1,
|
||||
output_dir = str(tmp_path),
|
||||
packing = True,
|
||||
|
|
@ -974,7 +976,12 @@ def test_enable_sample_packing():
|
|||
|
||||
|
||||
def test_enable_sample_packing_trl_collator(tmp_path):
|
||||
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
||||
if torch.cuda.is_available():
|
||||
device = torch.device("cuda")
|
||||
elif torch.xpu.is_available():
|
||||
device = torch.device("xpu")
|
||||
else:
|
||||
device = torch.device("cpu")
|
||||
model, _, trainer, _ = _build_packed_training_setup(tmp_path, device)
|
||||
|
||||
enable_sample_packing(model, trainer)
|
||||
|
|
@ -1030,7 +1037,12 @@ def test_enable_padding_free_metadata():
|
|||
|
||||
|
||||
def test_packing_sdpa(tmp_path):
|
||||
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
||||
if torch.cuda.is_available():
|
||||
device = torch.device("cuda")
|
||||
elif torch.xpu.is_available():
|
||||
device = torch.device("xpu")
|
||||
else:
|
||||
device = torch.device("cpu")
|
||||
model, batch, trainer, llama_mod = _build_packed_training_setup(tmp_path, device)
|
||||
|
||||
assert "packed_seq_lengths" in batch
|
||||
|
|
|
|||
|
|
@ -130,8 +130,14 @@ def _test_fake_quantizers_are_called(
|
|||
# Weight fake quantizers must always be called.
|
||||
assert child.weight_fake_quantizer.count == 1
|
||||
|
||||
if torch.cuda.is_available():
|
||||
device = torch.device("cuda")
|
||||
elif torch.xpu.is_available():
|
||||
device = torch.device("xpu")
|
||||
else:
|
||||
pytest.skip("No GPU available")
|
||||
for k, v in example_inputs.items():
|
||||
example_inputs[k] = v.cuda()
|
||||
example_inputs[k] = v.to(device)
|
||||
model.apply(_swap_fake_quantizers)
|
||||
model(**example_inputs)
|
||||
model.apply(_assert_fake_quantizers_are_called)
|
||||
|
|
|
|||
|
|
@ -15,18 +15,20 @@ import pytest
|
|||
import torch
|
||||
|
||||
|
||||
def _has_real_cuda():
|
||||
try:
|
||||
torch.zeros(1).to("cuda")
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
def _has_real_gpu():
|
||||
for backend in ("cuda", "xpu"):
|
||||
try:
|
||||
torch.zeros(1).to(backend)
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
return False
|
||||
|
||||
|
||||
HAS_REAL_CUDA = _has_real_cuda()
|
||||
requires_cuda = pytest.mark.skipif(
|
||||
not HAS_REAL_CUDA,
|
||||
reason = "LlamaRotaryEmbedding builds per-device CUDA caches in __init__",
|
||||
HAS_REAL_GPU = _has_real_gpu()
|
||||
requires_gpu = pytest.mark.skipif(
|
||||
not HAS_REAL_GPU,
|
||||
reason = "LlamaRotaryEmbedding builds per-device caches in __init__ (needs CUDA or XPU)",
|
||||
)
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
|
|
@ -360,7 +362,7 @@ def _cos_at_position(rot, position):
|
|||
# --- Layer 3: CUDA behavioral guard (real instantiation needs a device) ---
|
||||
|
||||
|
||||
@requires_cuda
|
||||
@requires_gpu
|
||||
def test_constructor_applies_llama3_scaling():
|
||||
config = _make_config(LLAMA3_ROPE_SCALING)
|
||||
rot = _unsloth_rotary(config)
|
||||
|
|
@ -371,7 +373,7 @@ def test_constructor_applies_llama3_scaling():
|
|||
), "LlamaRotaryEmbedding built from a llama3 config produced unscaled inv_freq (issue #2405)."
|
||||
|
||||
|
||||
@requires_cuda
|
||||
@requires_gpu
|
||||
def test_constructor_unscaled_config_uses_vanilla_inv_freq():
|
||||
rot = _unsloth_rotary(_make_config(None))
|
||||
got = rot.inv_freq.float().cpu()
|
||||
|
|
@ -381,7 +383,7 @@ def test_constructor_unscaled_config_uses_vanilla_inv_freq():
|
|||
), "LlamaRotaryEmbedding with no rope_scaling must use the vanilla inv_freq"
|
||||
|
||||
|
||||
@requires_cuda
|
||||
@requires_gpu
|
||||
def test_cos_cache_differs_between_scaled_and_unscaled_at_long_position():
|
||||
scaled = _unsloth_rotary(_make_config(LLAMA3_ROPE_SCALING))
|
||||
unscaled = _unsloth_rotary(_make_config(None))
|
||||
|
|
@ -397,7 +399,7 @@ def test_cos_cache_differs_between_scaled_and_unscaled_at_long_position():
|
|||
)
|
||||
|
||||
|
||||
@requires_cuda
|
||||
@requires_gpu
|
||||
def test_extended_cache_keeps_scaling_after_growth():
|
||||
scaled = _unsloth_rotary(_make_config(LLAMA3_ROPE_SCALING))
|
||||
# Grow past the initial cache size (mirrors long-context decode).
|
||||
|
|
@ -456,7 +458,7 @@ def _build_longrope_rotary():
|
|||
return rot, config
|
||||
|
||||
|
||||
@requires_cuda
|
||||
@requires_gpu
|
||||
@pytest.mark.parametrize(
|
||||
"build", [_build_llama3_rotary, _build_longrope_rotary], ids = ["llama3", "longrope"]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -374,7 +374,15 @@ elif DEVICE_TYPE == "hip":
|
|||
# NO-OP for rocm device
|
||||
pass
|
||||
elif DEVICE_TYPE == "xpu":
|
||||
import bitsandbytes as bnb
|
||||
# Same degradation as the cuda branch above: no bnb means no 4bit, not a
|
||||
# failed `import unsloth`.
|
||||
try:
|
||||
import bitsandbytes as bnb
|
||||
except Exception:
|
||||
print(
|
||||
"Unsloth: `bitsandbytes` is not installed - 4bit QLoRA unallowed, but 16bit and full finetuning works!"
|
||||
)
|
||||
bnb = None
|
||||
|
||||
# TODO: check triton for intel installed properly.
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -216,19 +216,32 @@ class RawTextDataLoader:
|
|||
|
||||
def _read_file_by_format(self, file_path, file_format):
|
||||
"""Read file content based on detected format."""
|
||||
with open(file_path, "r", encoding = "utf-8") as f:
|
||||
# utf-8-sig: Windows tooling (PowerShell's Out-File, Excel's "CSV UTF-8") prepends
|
||||
# a BOM that plain utf-8 keeps as a leading character. Without a BOM it decodes
|
||||
# exactly like utf-8.
|
||||
with open(file_path, "r", encoding = "utf-8-sig") as f:
|
||||
if file_format == "plain_text" or file_format == "markdown":
|
||||
return f.read()
|
||||
elif file_format == "json_lines":
|
||||
lines = []
|
||||
for line in f:
|
||||
if Path(file_path).suffix.lower() == ".json":
|
||||
# A .json file is a single JSON document (commonly a list
|
||||
# of records), so parsing it per line drops the whole file.
|
||||
try:
|
||||
data = json.loads(line.strip())
|
||||
text = self._extract_text_from_json(data)
|
||||
if text:
|
||||
lines.append(text)
|
||||
parsed = json.load(f)
|
||||
records = parsed if isinstance(parsed, list) else [parsed]
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
# Some files carry JSON Lines under a .json name.
|
||||
f.seek(0)
|
||||
records = self._iter_json_lines(f)
|
||||
else:
|
||||
# A .jsonl file is one JSON value per line: stay streaming so
|
||||
# a large file is never held in memory all at once.
|
||||
records = self._iter_json_lines(f)
|
||||
lines = []
|
||||
for data in records:
|
||||
text = self._extract_text_from_json(data)
|
||||
if text:
|
||||
lines.append(text)
|
||||
return "\n\n".join(lines)
|
||||
elif file_format == "csv_text_column":
|
||||
reader = csv.DictReader(f)
|
||||
|
|
@ -244,6 +257,17 @@ class RawTextDataLoader:
|
|||
_TEXT_FIELDS = ("text", "content", "message", "body", "description", "prompt")
|
||||
_TEXT_COLUMNS = _TEXT_FIELDS
|
||||
|
||||
def _iter_json_lines(self, handle):
|
||||
"""Yield one parsed JSON value per line, skipping blank and malformed lines."""
|
||||
for line in handle:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
yield json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
def _extract_text_from_json(self, data):
|
||||
"""Extract text from JSON object using common field names."""
|
||||
# Skip non-object lines (str/list/number): `field in data` would be a
|
||||
|
|
|
|||
|
|
@ -117,6 +117,17 @@ DEVICE_COUNT: int = get_device_count()
|
|||
ALLOW_PREQUANTIZED_MODELS: bool = True
|
||||
# HSA_STATUS_ERROR_EXCEPTION checks - sometimes AMD fails for BnB
|
||||
ALLOW_BITSANDBYTES: bool = True
|
||||
# Unusable bitsandbytes on any backend, not just hip: clear the flags the loader
|
||||
# reads before it selects a 4bit checkpoint. Same guarded import the fallbacks in
|
||||
# _gpu_init.py and kernels/utils.py use rather than a find_spec probe, so an
|
||||
# installed-but-broken wheel (missing .so, wrong ROCm/CUDA build) is treated as
|
||||
# unavailable by all three, not only by the ones that import it.
|
||||
try:
|
||||
import bitsandbytes as _bnb_probe
|
||||
del _bnb_probe
|
||||
except Exception:
|
||||
ALLOW_PREQUANTIZED_MODELS = False
|
||||
ALLOW_BITSANDBYTES = False
|
||||
# gfx906 (MI50 / Radeon VII / Vega 20): Dynamo/Inductor codegen is broken on this
|
||||
# legacy GCN arch (ROCm dropped it after 6.3) - compiled graphs crash or miscompile
|
||||
# while the eager path trains fine. Default compile off; setdefault so a user
|
||||
|
|
|
|||
|
|
@ -133,11 +133,28 @@ def calculate_settings(
|
|||
|
||||
|
||||
HAS_CUDA_STREAM = False
|
||||
import bitsandbytes as bnb
|
||||
try:
|
||||
import bitsandbytes as bnb
|
||||
except Exception:
|
||||
# device_type.py already degrades to 16bit/full finetuning when bnb is missing
|
||||
# (e.g. gfx906, whose generic wheel has no kernels). Keep the import working and
|
||||
# fail only if a 4bit path is actually entered.
|
||||
bnb = None
|
||||
|
||||
# https://github.com/bitsandbytes-foundation/bitsandbytes/pull/1330/files
|
||||
HAS_CUDA_STREAM = Version(bnb.__version__) > Version("0.43.3")
|
||||
get_ptr = bnb.functional.get_ptr
|
||||
|
||||
def _bnb_required(*args, **kwargs):
|
||||
raise RuntimeError(
|
||||
"Unsloth: 4bit QLoRA needs `bitsandbytes`, which is not installed. "
|
||||
"16bit LoRA and full finetuning work without it."
|
||||
)
|
||||
|
||||
|
||||
if bnb is not None:
|
||||
# https://github.com/bitsandbytes-foundation/bitsandbytes/pull/1330/files
|
||||
HAS_CUDA_STREAM = Version(bnb.__version__) > Version("0.43.3")
|
||||
get_ptr = bnb.functional.get_ptr
|
||||
else:
|
||||
get_ptr = _bnb_required
|
||||
|
||||
if DEVICE_TYPE == "xpu":
|
||||
HAS_XPU_STREAM = True
|
||||
|
|
@ -235,18 +252,25 @@ else:
|
|||
# Bitsandbytes operations
|
||||
ctypes_c_int = ctypes.c_int
|
||||
ctypes_c_int32 = ctypes.c_int32
|
||||
cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32
|
||||
cdequantize_blockwise_fp16_nf4 = bnb.functional.lib.cdequantize_blockwise_fp16_nf4
|
||||
cdequantize_blockwise_bf16_nf4 = bnb.functional.lib.cdequantize_blockwise_bf16_nf4
|
||||
|
||||
if DEVICE_TYPE == "xpu":
|
||||
# https://github.com/bitsandbytes-foundation/bitsandbytes/blob/c3b8de268fdb55a88f92feada23fc811a1e6877a/bitsandbytes/backends/xpu/ops.py#L115
|
||||
# for xpu, inference gemv using above link
|
||||
cgemm_4bit_inference_naive_fp16 = bnb.functional.lib.cgemv_4bit_inference_fp16
|
||||
cgemm_4bit_inference_naive_bf16 = bnb.functional.lib.cgemv_4bit_inference_bf16
|
||||
if bnb is None:
|
||||
cdequantize_blockwise_fp32 = _bnb_required
|
||||
cdequantize_blockwise_fp16_nf4 = _bnb_required
|
||||
cdequantize_blockwise_bf16_nf4 = _bnb_required
|
||||
cgemm_4bit_inference_naive_fp16 = _bnb_required
|
||||
cgemm_4bit_inference_naive_bf16 = _bnb_required
|
||||
else:
|
||||
cgemm_4bit_inference_naive_fp16 = bnb.functional.lib.cgemm_4bit_inference_naive_fp16
|
||||
cgemm_4bit_inference_naive_bf16 = bnb.functional.lib.cgemm_4bit_inference_naive_bf16
|
||||
cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32
|
||||
cdequantize_blockwise_fp16_nf4 = bnb.functional.lib.cdequantize_blockwise_fp16_nf4
|
||||
cdequantize_blockwise_bf16_nf4 = bnb.functional.lib.cdequantize_blockwise_bf16_nf4
|
||||
|
||||
if DEVICE_TYPE == "xpu":
|
||||
# https://github.com/bitsandbytes-foundation/bitsandbytes/blob/c3b8de268fdb55a88f92feada23fc811a1e6877a/bitsandbytes/backends/xpu/ops.py#L115
|
||||
# for xpu, inference gemv using above link
|
||||
cgemm_4bit_inference_naive_fp16 = bnb.functional.lib.cgemv_4bit_inference_fp16
|
||||
cgemm_4bit_inference_naive_bf16 = bnb.functional.lib.cgemv_4bit_inference_bf16
|
||||
else:
|
||||
cgemm_4bit_inference_naive_fp16 = bnb.functional.lib.cgemm_4bit_inference_naive_fp16
|
||||
cgemm_4bit_inference_naive_bf16 = bnb.functional.lib.cgemm_4bit_inference_naive_bf16
|
||||
|
||||
|
||||
torch_device_stream = (
|
||||
|
|
|
|||
|
|
@ -31,8 +31,20 @@ from .llama import (
|
|||
LlamaLinearScalingRotaryEmbedding,
|
||||
)
|
||||
from .mistral import *
|
||||
from bitsandbytes.nn import Linear4bit as Bnb_Linear4bit
|
||||
from peft.tuners.lora import Linear4bit as Peft_Linear4bit
|
||||
|
||||
# Without bnb, peft stops exporting its 4bit LoRA layer too. Both names only feed
|
||||
# isinstance checks, so placeholders nothing can match are exact stand-ins.
|
||||
try:
|
||||
from bitsandbytes.nn import Linear4bit as Bnb_Linear4bit
|
||||
from peft.tuners.lora import Linear4bit as Peft_Linear4bit
|
||||
except Exception:
|
||||
|
||||
class Bnb_Linear4bit:
|
||||
pass
|
||||
|
||||
class Peft_Linear4bit:
|
||||
pass
|
||||
|
||||
|
||||
try:
|
||||
from transformers.models.granite.modeling_granite import (
|
||||
|
|
|
|||
|
|
@ -472,13 +472,42 @@ class FastLanguageModel(FastLlamaModel):
|
|||
fast_inference = False
|
||||
break
|
||||
|
||||
# Check if 4bit is allowed specifically for AMD
|
||||
if not ALLOW_BITSANDBYTES and not use_exact_model_name:
|
||||
if load_in_4bit or load_in_8bit or model_name.lower().endswith("-bnb-4bit"):
|
||||
print(
|
||||
"Unsloth: AMD currently is not stable with 4bit bitsandbytes. Disabling for now."
|
||||
# bitsandbytes unusable (absent, or unstable as on some AMD stacks). This is
|
||||
# a capability check, so it is not gated on use_exact_model_name: that only
|
||||
# suppresses repo-name remapping and cannot make bitsandbytes available.
|
||||
if not ALLOW_BITSANDBYTES:
|
||||
# A user-supplied config sets load_in_4bit/8bit above and is forwarded
|
||||
# in kwargs, so clearing the flags alone still rebuilds the bnb
|
||||
# quantizer downstream. Only drop it when it asks for bnb: a GPTQ /
|
||||
# AWQ / fp8 / torchao config must pass through untouched.
|
||||
_quant_cfg = kwargs.get("quantization_config", None)
|
||||
if isinstance(_quant_cfg, dict):
|
||||
_wants_bnb = bool(
|
||||
_quant_cfg.get("load_in_4bit", False) or _quant_cfg.get("load_in_8bit", False)
|
||||
)
|
||||
elif _quant_cfg is not None:
|
||||
_wants_bnb = bool(
|
||||
getattr(_quant_cfg, "load_in_4bit", False)
|
||||
or getattr(_quant_cfg, "load_in_8bit", False)
|
||||
)
|
||||
else:
|
||||
_wants_bnb = False
|
||||
if (
|
||||
load_in_4bit
|
||||
or load_in_8bit
|
||||
or _wants_bnb
|
||||
or model_name.lower().endswith("-bnb-4bit")
|
||||
):
|
||||
print(
|
||||
"Unsloth: `bitsandbytes` is unavailable here - disabling 4bit/8bit. "
|
||||
"16bit LoRA and full finetuning still work."
|
||||
)
|
||||
# 8bit is bitsandbytes too: leaving either set sends the request on to
|
||||
# Transformers, which builds the bnb quantizer and fails there.
|
||||
load_in_4bit = False
|
||||
load_in_8bit = False
|
||||
if _wants_bnb:
|
||||
kwargs.pop("quantization_config", None)
|
||||
|
||||
# Find FP8, BnB 4bit, other mapped names
|
||||
old_model_name = model_name
|
||||
|
|
@ -1102,7 +1131,13 @@ class FastModel(FastBaseModel):
|
|||
assert load_in_fp8 in (True, False, "block")
|
||||
|
||||
patch_compiled_autograd()
|
||||
patch_compiling_bitsandbytes()
|
||||
# Same best-effort wrapper as the FastLanguageModel path: unsloth_zoo's
|
||||
# patch imports bitsandbytes unconditionally, so on a host without it this
|
||||
# raised before the capability fallback below could take the 16bit path.
|
||||
try:
|
||||
patch_compiling_bitsandbytes()
|
||||
except Exception as e:
|
||||
print(f"Unsloth: Could not patch bitsandbytes for torch.compile - {e}")
|
||||
|
||||
if full_finetuning and (load_in_4bit or load_in_8bit):
|
||||
print(
|
||||
|
|
@ -1113,6 +1148,43 @@ class FastModel(FastBaseModel):
|
|||
load_in_fp8 = False
|
||||
load_in_16bit = False
|
||||
|
||||
# bitsandbytes unusable (absent, or unstable as on some AMD stacks). This is
|
||||
# a capability check, so it is not gated on use_exact_model_name: that only
|
||||
# suppresses repo-name remapping and cannot make bitsandbytes available.
|
||||
if not ALLOW_BITSANDBYTES:
|
||||
# A user-supplied config sets load_in_4bit/8bit above and is forwarded
|
||||
# in kwargs, so clearing the flags alone still rebuilds the bnb
|
||||
# quantizer downstream. Only drop it when it asks for bnb: a GPTQ /
|
||||
# AWQ / fp8 / torchao config must pass through untouched.
|
||||
_quant_cfg = kwargs.get("quantization_config", None)
|
||||
if isinstance(_quant_cfg, dict):
|
||||
_wants_bnb = bool(
|
||||
_quant_cfg.get("load_in_4bit", False) or _quant_cfg.get("load_in_8bit", False)
|
||||
)
|
||||
elif _quant_cfg is not None:
|
||||
_wants_bnb = bool(
|
||||
getattr(_quant_cfg, "load_in_4bit", False)
|
||||
or getattr(_quant_cfg, "load_in_8bit", False)
|
||||
)
|
||||
else:
|
||||
_wants_bnb = False
|
||||
if (
|
||||
load_in_4bit
|
||||
or load_in_8bit
|
||||
or _wants_bnb
|
||||
or model_name.lower().endswith("-bnb-4bit")
|
||||
):
|
||||
print(
|
||||
"Unsloth: `bitsandbytes` is unavailable here - disabling 4bit/8bit. "
|
||||
"16bit LoRA and full finetuning still work."
|
||||
)
|
||||
# 8bit is bitsandbytes too: leaving either set sends the request on to
|
||||
# Transformers, which builds the bnb quantizer and fails there.
|
||||
load_in_4bit = False
|
||||
load_in_8bit = False
|
||||
if _wants_bnb:
|
||||
kwargs.pop("quantization_config", None)
|
||||
|
||||
if (
|
||||
int(load_in_4bit) + int(load_in_8bit) + int(load_in_16bit) + int(load_in_fp8 != False)
|
||||
>= 2
|
||||
|
|
@ -1142,14 +1214,6 @@ class FastModel(FastBaseModel):
|
|||
if is_dist:
|
||||
device_map = distributed_device_map
|
||||
|
||||
# Check if 4bit is allowed specifically for AMD
|
||||
if not ALLOW_BITSANDBYTES and not use_exact_model_name:
|
||||
if load_in_4bit or load_in_8bit or model_name.lower().endswith("-bnb-4bit"):
|
||||
print(
|
||||
"Unsloth: AMD currently is not stable with 4bit bitsandbytes. Disabling for now."
|
||||
)
|
||||
load_in_4bit = False
|
||||
|
||||
if fast_inference:
|
||||
if importlib.util.find_spec("vllm") is None:
|
||||
raise ImportError(
|
||||
|
|
|
|||
|
|
@ -32,8 +32,20 @@ except ImportError:
|
|||
import sys
|
||||
IS_WINDOWS = sys.platform == "win32"
|
||||
LLAMA_CPP_DEFAULT_DIR = "llama.cpp"
|
||||
from bitsandbytes.nn import Linear4bit as Bnb_Linear4bit
|
||||
from peft.tuners.lora import Linear4bit as Peft_Linear4bit
|
||||
# Without bnb, peft stops exporting its 4bit LoRA layer too. Both names only feed
|
||||
# isinstance checks, so placeholders nothing can match are exact stand-ins.
|
||||
try:
|
||||
from bitsandbytes.nn import Linear4bit as Bnb_Linear4bit
|
||||
from peft.tuners.lora import Linear4bit as Peft_Linear4bit
|
||||
except Exception:
|
||||
|
||||
class Bnb_Linear4bit:
|
||||
pass
|
||||
|
||||
class Peft_Linear4bit:
|
||||
pass
|
||||
|
||||
|
||||
from peft.tuners.lora import Linear as Peft_Linear
|
||||
from typing import Optional, Callable, Union, List
|
||||
import sys
|
||||
|
|
@ -3843,10 +3855,10 @@ from .models.loader_utils import (
|
|||
_tokenizer_cache_dir,
|
||||
_tokenizer_wants_local_only,
|
||||
)
|
||||
from unsloth_zoo.saving_utils import (
|
||||
merge_and_overwrite_lora,
|
||||
prepare_saving,
|
||||
)
|
||||
|
||||
# Imported lazily at the two call sites below: a zoo older than the one that made
|
||||
# its own bitsandbytes import optional would otherwise break `import unsloth` on a
|
||||
# host without bnb, which is the whole point of the guards above.
|
||||
from unsloth_zoo.llama_cpp import (
|
||||
install_llama_cpp,
|
||||
convert_to_gguf as _convert_to_gguf,
|
||||
|
|
@ -4094,6 +4106,8 @@ def save_to_gguf_generic(
|
|||
quantization_type = quantization_type,
|
||||
)
|
||||
if repo_id is not None:
|
||||
from unsloth_zoo.saving_utils import prepare_saving
|
||||
|
||||
prepare_saving(
|
||||
model,
|
||||
repo_id,
|
||||
|
|
@ -4225,6 +4239,7 @@ def unsloth_generic_save(
|
|||
print(f"Unsloth: Model saved successfully to '{save_directory}'")
|
||||
else:
|
||||
_prewarm_base_model_hub_cache(model, save_method = save_method, token = token)
|
||||
from unsloth_zoo.saving_utils import merge_and_overwrite_lora
|
||||
merge_and_overwrite_lora(
|
||||
get_model_name,
|
||||
model = model,
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Add a link
Reference in a new issue