diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index 1758efe515..5aca8e3996 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -394,6 +394,9 @@ class LoadResponse(BaseModel): is_vision: bool = Field(False, description = "Whether model is a vision model") is_lora: bool = Field(False, description = "Whether model is a LoRA adapter") is_gguf: bool = Field(False, description = "Whether model is a GGUF model (llama.cpp)") + is_local_model: bool = Field( + False, description = "Whether the loaded model came from a local filesystem path" + ) is_diffusion: bool = Field( False, description = "Whether model is a block-diffusion model (DiffusionGemma)" ) @@ -558,6 +561,9 @@ class InferenceStatusResponse(BaseModel): ) is_vision: bool = Field(False, description = "Whether the active model is a vision model") is_gguf: bool = Field(False, description = "Whether the active model is a GGUF model (llama.cpp)") + is_local_model: bool = Field( + False, description = "Whether the active model came from a local filesystem path" + ) is_diffusion: bool = Field( False, description = "Whether the active model is a block-diffusion model (DiffusionGemma)" ) diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index f4b8a94af5..91a7022398 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -1018,6 +1018,7 @@ try: ) from core.inference.tensor_fallback import load_with_tensor_fallback from utils.models import ModelConfig + from utils.paths import is_local_path from utils.inference import load_inference_config from utils.models.model_config import ( _local_gguf_companion_search_root, @@ -1058,6 +1059,7 @@ except ImportError: ) from core.inference.tensor_fallback import load_with_tensor_fallback from utils.models import ModelConfig + from utils.paths import is_local_path from utils.inference import load_inference_config from utils.models.model_config import ( _local_gguf_companion_search_root, @@ -4420,6 +4422,8 @@ async def _load_model_impl( is_vision = llama_backend._is_vision, is_lora = False, is_gguf = True, + is_local_model = native_grant_backed + or is_local_path(llama_backend.model_identifier), is_diffusion = llama_backend.is_diffusion, is_audio = _gguf_is_audio, audio_type = _gguf_audio, @@ -4478,6 +4482,7 @@ async def _load_model_impl( is_vision = _model_info.get("is_vision", False), is_lora = _model_info.get("is_lora", False), is_gguf = False, + is_local_model = native_grant_backed or is_local_path(backend.active_model_name), is_audio = _model_info.get("is_audio", False), audio_type = _model_info.get("audio_type"), has_audio_input = _model_info.get("has_audio_input", False), @@ -4798,6 +4803,7 @@ async def _load_model_impl( is_vision = llama_backend.is_vision, is_lora = False, is_gguf = True, + is_local_model = config.is_local, is_diffusion = llama_backend.is_diffusion, is_audio = _gguf_is_audio, audio_type = _gguf_audio, @@ -4938,6 +4944,7 @@ async def _load_model_impl( is_vision = config.is_vision, is_lora = config.is_lora, is_gguf = False, + is_local_model = config.is_local, is_audio = config.is_audio, audio_type = config.audio_type, has_audio_input = config.has_audio_input, @@ -5918,6 +5925,7 @@ async def get_status(current_subject: str = Depends(get_current_subject)): model_identifier = None if _native_grant_backed else _model_id, is_vision = llama_backend.is_vision, is_gguf = True, + is_local_model = _native_grant_backed or bool(_model_id and is_local_path(_model_id)), is_diffusion = llama_backend.is_diffusion, gguf_variant = llama_backend.hf_variant, is_audio = getattr(llama_backend, "_is_audio", False), @@ -5989,6 +5997,9 @@ async def get_status(current_subject: str = Depends(get_current_subject)): model_identifier = backend.active_model_name, is_vision = is_vision, is_gguf = False, + is_local_model = bool( + backend.active_model_name and is_local_path(backend.active_model_name) + ), is_audio = is_audio, audio_type = audio_type, has_audio_input = has_audio_input, diff --git a/studio/backend/tests/test_mtp_drafter_companion.py b/studio/backend/tests/test_mtp_drafter_companion.py index 5c1e7fb035..a3ac2e5071 100644 --- a/studio/backend/tests/test_mtp_drafter_companion.py +++ b/studio/backend/tests/test_mtp_drafter_companion.py @@ -29,7 +29,9 @@ from hub.utils.gguf_plan import ( preferred_mtp_sibling, ) from utils.models.model_config import ( + ModelConfig, _is_mtp_drafter, + _local_gguf_companion_search_root, detect_gguf_model, detect_mtp_file, extract_model_size_b, @@ -212,6 +214,36 @@ def test_detect_mtp_file_search_root(tmp_path): assert found is not None and found.endswith("mtp-gemma-4-12b-it.gguf") +def test_quant_directory_selection_finds_repo_root_mtp(tmp_path): + quant_dir = tmp_path / "Q4_0" + quant_dir.mkdir() + weight = quant_dir / "gemma-4-E4B-it-qat-Q4_0.gguf" + weight.write_bytes(b"x") + mtp_dir = tmp_path / "MTP" + mtp_dir.mkdir() + drafter = mtp_dir / "mtp-gemma-4-E4B-it-Q4_0.gguf" + drafter.write_bytes(b"x") + + search_root = _local_gguf_companion_search_root(str(quant_dir), str(weight)) + assert Path(search_root).resolve() == tmp_path.resolve() + config = ModelConfig.from_identifier(str(quant_dir)) + assert config.is_local + assert config.gguf_file == str(weight.resolve()) + assert config.gguf_mtp_file == str(drafter.resolve()) + + +def test_bare_relative_gguf_directory_is_local_source(tmp_path, monkeypatch): + model_dir = tmp_path / "outputs" / "gemma" + model_dir.mkdir(parents = True) + weight = model_dir / "gemma-4-E4B-it-qat-Q4_0.gguf" + weight.write_bytes(b"x") + monkeypatch.chdir(tmp_path) + + config = ModelConfig.from_identifier("outputs/gemma") + assert config.is_local + assert config.gguf_file == str(weight.resolve()) + + def test_detect_mtp_file_falls_back_to_new_scheme_subdir(tmp_path): weight = tmp_path / "gemma-4-E4B-it-qat-Q4_0.gguf" weight.write_bytes(b"x") diff --git a/studio/backend/tests/test_native_gguf_companion.py b/studio/backend/tests/test_native_gguf_companion.py index 9ed260aa90..33e3c70dbb 100644 --- a/studio/backend/tests/test_native_gguf_companion.py +++ b/studio/backend/tests/test_native_gguf_companion.py @@ -99,6 +99,25 @@ def test_reload_dedup_finds_repo_root_mtp_companion(tmp_path, monkeypatch): assert _request_matches_loaded_settings(request, backend) +def test_reload_dedup_matches_quant_directory_selection(tmp_path, monkeypatch): + quant_dir = tmp_path / "Q4_0" + quant_dir.mkdir() + weight = quant_dir / "model.gguf" + weight.write_bytes(b"model") + companion_dir = tmp_path / "MTP" + companion_dir.mkdir() + companion = companion_dir / "mtp-model.gguf" + companion.write_bytes(b"draft") + + monkeypatch.setattr(LlamaCppBackend, "_kill_orphaned_servers", staticmethod(lambda: 0)) + backend = LlamaCppBackend() + backend._gguf_path = str(weight) + backend._mtp_draft_path = str(companion) + + request = LoadRequest(model_path = str(quant_dir), gguf_variant = "Q4_0") + assert _request_matches_loaded_settings(request, backend) + + def test_native_vision_companion_rejects_mtp_directory(tmp_path): weight, companion = _write_pair(tmp_path, "MTP") with pytest.raises(HTTPException, match = "must live next to"): diff --git a/studio/backend/utils/models/model_config.py b/studio/backend/utils/models/model_config.py index 0d66515182..34c0acda64 100644 --- a/studio/backend/utils/models/model_config.py +++ b/studio/backend/utils/models/model_config.py @@ -1765,13 +1765,6 @@ def _local_gguf_companion_search_root(selected_path: str, gguf_file: str) -> str selected = Path(selected_path) gguf_path = Path(gguf_file) - if selected.suffix.lower() != ".gguf": - return selected_path - - gguf_dir = gguf_path.parent - if not gguf_dir.name: - return str(gguf_dir) - quant_dir_re = ( r"(UD-)?(" r"MXFP[0-9]+(?:_[A-Z0-9]+)*" @@ -1783,9 +1776,12 @@ def _local_gguf_companion_search_root(selected_path: str, gguf_file: str) -> str r"|BF16|F16|F32" r")" ) - if re.fullmatch(quant_dir_re, gguf_dir.name, re.IGNORECASE): - return str(gguf_dir.parent) - return str(gguf_dir) + search_dir = gguf_path.parent if selected.suffix.lower() == ".gguf" else selected + if not search_dir.name: + return str(search_dir) + if re.fullmatch(quant_dir_re, search_dir.name, re.IGNORECASE): + return str(search_dir.parent) + return str(search_dir) def _iter_hf_cache_snapshots(repo_id: str, cache_dir: Optional[str | Path] = None): diff --git a/studio/frontend/src/features/chat/api/chat-adapter.ts b/studio/frontend/src/features/chat/api/chat-adapter.ts index d4861e8a3a..822b8627be 100644 --- a/studio/frontend/src/features/chat/api/chat-adapter.ts +++ b/studio/frontend/src/features/chat/api/chat-adapter.ts @@ -1704,6 +1704,7 @@ async function autoLoadSmallestModel(): Promise<{ customContextLength: config.customContextLength, loadedIsMultimodal: isMultimodalResponse(loadResp), loadedIsDiffusion: loadResp.is_diffusion ?? false, + activeModelIsLocal: loadResp.is_local_model ?? false, ...resolveLoadedSpeculativeSettings(loadResp), }); } else { @@ -1729,6 +1730,7 @@ async function autoLoadSmallestModel(): Promise<{ ...resolveLoadedSpeculativeSettings(loadResp), loadedIsMultimodal: isMultimodalResponse(loadResp), loadedIsDiffusion: loadResp.is_diffusion ?? false, + activeModelIsLocal: loadResp.is_local_model ?? false, }); } if (!(loadResp.is_lora ?? false)) { @@ -2007,6 +2009,7 @@ async function autoLoadSmallestModel(): Promise<{ defaultChatTemplate: loadResp.chat_template ?? null, chatTemplateOverride: null, loadedIsMultimodal: isMultimodalResponse(loadResp), + activeModelIsLocal: loadResp.is_local_model ?? false, ...resolveLoadedSpeculativeSettings(loadResp), }); recordLastLocalModelLoad({ diff --git a/studio/frontend/src/features/chat/chat-settings-sheet.tsx b/studio/frontend/src/features/chat/chat-settings-sheet.tsx index a773db63b6..c9d59d1b14 100644 --- a/studio/frontend/src/features/chat/chat-settings-sheet.tsx +++ b/studio/frontend/src/features/chat/chat-settings-sheet.tsx @@ -384,6 +384,9 @@ export function ChatSettingsPanel({ const activeNativePathToken = useChatRuntimeStore( (s) => s.activeNativePathToken, ); + const activeModelIsLocal = useChatRuntimeStore( + (s) => s.activeModelIsLocal, + ); const ggufContextLength = useChatRuntimeStore((s) => s.ggufContextLength); // Direct-file / custom-folder GGUFs load without a variant label but still // report a GGUF context, so detect them via the context and the checkpoint @@ -395,7 +398,8 @@ export function ChatSettingsPanel({ (currentCheckpoint?.toLowerCase().endsWith(".gguf") ?? false); const isLocalGguf = isGguf && - (activeNativePathToken != null || + (activeModelIsLocal || + activeNativePathToken != null || isLocalModelPath(currentCheckpoint ?? "") || (currentCheckpoint?.toLowerCase().endsWith(".gguf") ?? false)); const ggufMaxContextLength = useChatRuntimeStore( diff --git a/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts b/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts index 76a310ac33..000eda6b62 100644 --- a/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts +++ b/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts @@ -319,6 +319,7 @@ async function syncInferenceStatusToStore(options?: { modelRequiresTrustRemoteCode: false, loadedIsMultimodal: false, loadedIsDiffusion: false, + activeModelIsLocal: false, }); } } catch (error) { @@ -1004,6 +1005,7 @@ export function useChatModelRuntime() { loadedChatTemplateOverride: effectiveChatTemplateOverride, loadedIsMultimodal: isMultimodalResponse(loadResponse), loadedIsDiffusion: loadResponse.is_diffusion ?? false, + activeModelIsLocal: loadResponse.is_local_model ?? false, activeNativePathToken: nativePathToken ?? null, activeNativePathExpiresAtMs: nativePathToken ? nativePathExpiresAtMs @@ -1113,6 +1115,7 @@ export function useChatModelRuntime() { rollbackResponse.speculative_type, ); useChatRuntimeStore.setState({ + activeModelIsLocal: rollbackResponse.is_local_model ?? false, activeNativePathToken: previousActiveNativePathToken ?? null, // Restore the previous token's lease together with the token so a // rollback never pairs restored token A with failed load B's expiry. diff --git a/studio/frontend/src/features/chat/lib/apply-inference-status-to-store.ts b/studio/frontend/src/features/chat/lib/apply-inference-status-to-store.ts index f85ff3246b..31ccd3721c 100644 --- a/studio/frontend/src/features/chat/lib/apply-inference-status-to-store.ts +++ b/studio/frontend/src/features/chat/lib/apply-inference-status-to-store.ts @@ -289,6 +289,7 @@ export function applyActiveModelStatusToStore( defaultChatTemplate: nextDefaultChatTemplate, loadedIsMultimodal: isMultimodalResponse(status), loadedIsDiffusion: status.is_diffusion ?? false, + activeModelIsLocal: status.is_local_model ?? false, specFallbackReason: status.spec_fallback_reason ?? null, // The spec / KV seeds share the GPU-fields reseed mechanism below: a // non-GGUF status leaves their loaded baselines null, so the "unseeded" diff --git a/studio/frontend/src/features/chat/shared-composer.tsx b/studio/frontend/src/features/chat/shared-composer.tsx index 3f89f70ef1..28d9e6673b 100644 --- a/studio/frontend/src/features/chat/shared-composer.tsx +++ b/studio/frontend/src/features/chat/shared-composer.tsx @@ -1226,6 +1226,7 @@ export function SharedComposer({ // GPU fields on every load path so the gate can't read stale. loadedIsDiffusion: resp.is_diffusion ?? false, loadedIsMultimodal: isMultimodalResponse(resp), + activeModelIsLocal: resp.is_local_model ?? false, // Record the context this pane loaded with (like the single-model path) // so when it becomes the active model, the UI and later reload/save use // its context, not the previous/default one. diff --git a/studio/frontend/src/features/chat/stores/chat-runtime-store.ts b/studio/frontend/src/features/chat/stores/chat-runtime-store.ts index 42359b8f7f..9409945072 100644 --- a/studio/frontend/src/features/chat/stores/chat-runtime-store.ts +++ b/studio/frontend/src/features/chat/stores/chat-runtime-store.ts @@ -735,6 +735,8 @@ type ChatRuntimeStore = { // lets the attach gates flag a failed load vs "no model picked". lastModelLoadError: string | null; activeGgufVariant: string | null; + /** Whether the backend loaded the active model from a filesystem path. */ + activeModelIsLocal: boolean; ggufContextLength: number | null; ggufMaxContextLength: number | null; ggufNativeContextLength: number | null; @@ -1268,6 +1270,7 @@ export const useChatRuntimeStore = create((set, get) => ({ modelsError: null, lastModelLoadError: null, activeGgufVariant: null, + activeModelIsLocal: false, ggufContextLength: null, ggufMaxContextLength: null, ggufNativeContextLength: null, @@ -1555,6 +1558,7 @@ export const useChatRuntimeStore = create((set, get) => ({ checkpoint: "", }, activeGgufVariant: null, + activeModelIsLocal: false, activeNativePathToken: null, activeNativePathExpiresAtMs: null, ggufContextLength: null, diff --git a/studio/frontend/src/features/chat/types/api.ts b/studio/frontend/src/features/chat/types/api.ts index 6c3e919efe..e5b4239947 100644 --- a/studio/frontend/src/features/chat/types/api.ts +++ b/studio/frontend/src/features/chat/types/api.ts @@ -152,6 +152,7 @@ export interface LoadModelResponse { is_vision: boolean; is_lora: boolean; is_gguf?: boolean; + is_local_model?: boolean; is_diffusion?: boolean; is_audio?: boolean; audio_type?: string | null; @@ -203,6 +204,7 @@ export interface InferenceStatusResponse { model_identifier?: string | null; is_vision: boolean; is_gguf?: boolean; + is_local_model?: boolean; is_diffusion?: boolean; gguf_variant?: string | null; is_audio?: boolean; diff --git a/tests/studio/test_model_picker_contracts.py b/tests/studio/test_model_picker_contracts.py index 1dc9761c12..c946a5de4a 100644 --- a/tests/studio/test_model_picker_contracts.py +++ b/tests/studio/test_model_picker_contracts.py @@ -26,6 +26,12 @@ def _read(rel: str) -> str: return path.read_text() +def _read_backend(rel: str) -> str: + path = WORKDIR / "studio" / "backend" / rel + assert path.exists(), f"missing backend source file: {path}" + return path.read_text() + + def test_models_api_sends_token_via_header_not_query(): """getModelConfig / checkVisionModel / checkEmbeddingModel must pass the HF token through hubTokenHeader, never as a ?hf_token= query param (which leaks @@ -326,10 +332,34 @@ def test_local_mtp_warning_covers_path_and_native_gguf_sources(): assert local assert "isGguf &&" in local.group(0) assert "activeNativePathToken" in local.group(0) + assert "activeModelIsLocal" in local.group(0) assert "isLocalModelPath" in local.group(0) assert "isLocalGguf" in src.split('specFallbackReason === "drafter_not_found"', 1)[1] +def test_local_mtp_warning_uses_backend_source_metadata(): + types = _read("features/chat/types/api.ts") + assert types.count("is_local_model?: boolean") >= 2 + + status = _read("features/chat/lib/apply-inference-status-to-store.ts") + assert "activeModelIsLocal: status.is_local_model ?? false" in status + + runtime = _read("features/chat/stores/chat-runtime-store.ts") + assert "activeModelIsLocal: boolean" in runtime + assert runtime.count("activeModelIsLocal: false") >= 2 + + load = _read("features/chat/hooks/use-chat-model-runtime.ts") + assert "activeModelIsLocal: loadResponse.is_local_model ?? false" in load + + models = _read_backend("models/inference.py") + assert models.count("is_local_model: bool = Field(") >= 2 + + route = _read_backend("routes/inference.py") + assert route.count("is_local_model = config.is_local") >= 2 + assert "is_local_model = _native_grant_backed" in route + assert "backend.active_model_name and is_local_path(backend.active_model_name)" in route + + def test_fixed_layer_gguf_pins_displayed_context(): """An already-loaded auto-fit GGUF saved with Manual fixed GPU layers must pin the shown context, so a later fresh load keeps the fitted placement