* feat: multi-source model discovery (HF default, legacy cache, LM Studio) * Fix multi-source model discovery bugs - Fix lmstudio_model_dirs: add ~/.lmstudio/models as default path, remove dead sys.platform branch, add dedup via seen set - Fix _setup_cache_env: preserve legacy HF cache env vars when the legacy hub directory exists and is non-empty - Fix _scan_lmstudio_dir: use absolute path for id field so is_local_path() returns True - Remove LM Studio dirs from allowed_roots (scanned unconditionally) - Replace bare except passes with logger.warning in legacy cache blocks - Fix delete_cached_model to search both default and legacy HF caches - Make lmstudio_dirs non-optional in TS interface (matches Python schema) - Exclude lmstudio source from trainable model filter - Remove unused import sys * Scan HF default cache alongside legacy and active caches When _setup_cache_env overrides HF_HUB_CACHE to the legacy Unsloth path, the standard HF default cache (~/.cache/huggingface/hub) was never scanned, hiding models downloaded before Unsloth Studio was installed. Add hf_default_cache_dir() and _all_hf_cache_scans() helper that deduplicates and scans all three HF cache locations (active, legacy, default). Used in list_local_models, list_cached_gguf, list_cached_models, and delete_cached_model. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: Daniel Han <danielhanchen@gmail.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
1228 lines
43 KiB
Python
1228 lines
43 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
|
|
|
"""
|
|
Model Management API routes
|
|
"""
|
|
|
|
import os
|
|
import sys
|
|
from pathlib import Path
|
|
from fastapi import APIRouter, Body, Depends, HTTPException, Query
|
|
from typing import List, Optional
|
|
import structlog
|
|
from loggers import get_logger
|
|
|
|
import re as _re
|
|
|
|
_VALID_REPO_ID = _re.compile(r"^[A-Za-z0-9._-]+/[A-Za-z0-9._-]+$")
|
|
|
|
|
|
def _is_valid_repo_id(repo_id: str) -> bool:
|
|
return bool(_VALID_REPO_ID.fullmatch(repo_id))
|
|
|
|
|
|
# Add backend directory to path
|
|
backend_path = Path(__file__).parent.parent.parent
|
|
if str(backend_path) not in sys.path:
|
|
sys.path.insert(0, str(backend_path))
|
|
|
|
from auth.authentication import get_current_subject
|
|
|
|
# Import backend functions
|
|
try:
|
|
from utils.models import (
|
|
scan_trained_loras,
|
|
scan_exported_models,
|
|
load_model_defaults,
|
|
get_base_model_from_lora,
|
|
is_vision_model,
|
|
is_embedding_model,
|
|
scan_checkpoints,
|
|
list_gguf_variants,
|
|
ModelConfig,
|
|
)
|
|
from utils.models.model_config import (
|
|
_pick_best_gguf,
|
|
_extract_quant_label,
|
|
is_audio_input_type,
|
|
)
|
|
from core.inference import get_inference_backend
|
|
from utils.paths import (
|
|
outputs_root,
|
|
exports_root,
|
|
resolve_output_dir,
|
|
resolve_export_dir,
|
|
)
|
|
except ImportError:
|
|
# Fallback: try to import from parent directory
|
|
parent_backend = backend_path.parent / "backend"
|
|
if str(parent_backend) not in sys.path:
|
|
sys.path.insert(0, str(parent_backend))
|
|
from utils.models import (
|
|
scan_trained_loras,
|
|
scan_exported_models,
|
|
load_model_defaults,
|
|
get_base_model_from_lora,
|
|
is_vision_model,
|
|
is_embedding_model,
|
|
scan_checkpoints,
|
|
list_gguf_variants,
|
|
ModelConfig,
|
|
)
|
|
from utils.models.model_config import (
|
|
_pick_best_gguf,
|
|
_extract_quant_label,
|
|
is_audio_input_type,
|
|
)
|
|
from core.inference import get_inference_backend
|
|
from utils.paths import (
|
|
outputs_root,
|
|
exports_root,
|
|
resolve_output_dir,
|
|
resolve_export_dir,
|
|
)
|
|
|
|
from models import (
|
|
CheckpointInfo,
|
|
CheckpointListResponse,
|
|
LocalModelInfo,
|
|
LocalModelListResponse,
|
|
ModelCheckpoints,
|
|
ModelDetails,
|
|
LoRAScanResponse,
|
|
LoRAInfo,
|
|
ModelListResponse,
|
|
)
|
|
from models.models import GgufVariantDetail, GgufVariantsResponse, ModelType
|
|
from models.responses import (
|
|
LoRABaseModelResponse,
|
|
VisionCheckResponse,
|
|
EmbeddingCheckResponse,
|
|
)
|
|
|
|
router = APIRouter()
|
|
logger = get_logger(__name__)
|
|
|
|
|
|
def derive_model_type(
|
|
is_vision: bool, audio_type: Optional[str], is_embedding: bool = False
|
|
) -> ModelType:
|
|
"""Collapse individual capability flags into a single model modality string."""
|
|
if is_embedding:
|
|
return "embeddings"
|
|
if audio_type is not None:
|
|
return "audio"
|
|
if is_vision:
|
|
return "vision"
|
|
return "text"
|
|
|
|
|
|
def _resolve_hf_cache_dir() -> Path:
|
|
"""Resolve local HF cache root used by hub downloads."""
|
|
try:
|
|
from huggingface_hub.constants import HF_HUB_CACHE
|
|
|
|
return Path(HF_HUB_CACHE)
|
|
except Exception:
|
|
return Path.home() / ".cache" / "huggingface" / "hub"
|
|
|
|
|
|
def _scan_models_dir(models_dir: Path) -> List[LocalModelInfo]:
|
|
if not models_dir.exists() or not models_dir.is_dir():
|
|
return []
|
|
|
|
found: List[LocalModelInfo] = []
|
|
for child in models_dir.iterdir():
|
|
if not child.is_dir():
|
|
continue
|
|
has_model_files = (
|
|
(child / "config.json").exists()
|
|
or (child / "adapter_config.json").exists()
|
|
or any(child.glob("*.safetensors"))
|
|
or any(child.glob("*.bin"))
|
|
or any(child.glob("*.gguf"))
|
|
)
|
|
if not has_model_files:
|
|
continue
|
|
try:
|
|
updated_at = child.stat().st_mtime
|
|
except OSError:
|
|
updated_at = None
|
|
found.append(
|
|
LocalModelInfo(
|
|
id = str(child),
|
|
display_name = child.name,
|
|
path = str(child),
|
|
source = "models_dir",
|
|
updated_at = updated_at,
|
|
),
|
|
)
|
|
# Also scan for standalone .gguf files directly in the models directory
|
|
for gguf_file in models_dir.glob("*.gguf"):
|
|
if gguf_file.is_file():
|
|
try:
|
|
updated_at = gguf_file.stat().st_mtime
|
|
except OSError:
|
|
updated_at = None
|
|
found.append(
|
|
LocalModelInfo(
|
|
id = str(gguf_file),
|
|
display_name = gguf_file.stem,
|
|
path = str(gguf_file),
|
|
source = "models_dir",
|
|
updated_at = updated_at,
|
|
),
|
|
)
|
|
|
|
return found
|
|
|
|
|
|
def _scan_hf_cache(cache_dir: Path) -> List[LocalModelInfo]:
|
|
if not cache_dir.exists() or not cache_dir.is_dir():
|
|
return []
|
|
|
|
found: List[LocalModelInfo] = []
|
|
for repo_dir in cache_dir.glob("models--*"):
|
|
if not repo_dir.is_dir():
|
|
continue
|
|
|
|
repo_name = repo_dir.name[len("models--") :]
|
|
if not repo_name:
|
|
continue
|
|
model_id = repo_name.replace("--", "/")
|
|
|
|
try:
|
|
updated_at = repo_dir.stat().st_mtime
|
|
except OSError:
|
|
updated_at = None
|
|
|
|
found.append(
|
|
LocalModelInfo(
|
|
id = model_id,
|
|
model_id = model_id,
|
|
display_name = model_id.split("/")[-1],
|
|
path = str(repo_dir),
|
|
source = "hf_cache",
|
|
updated_at = updated_at,
|
|
),
|
|
)
|
|
return found
|
|
|
|
|
|
def _scan_lmstudio_dir(lm_dir: Path) -> List[LocalModelInfo]:
|
|
"""Scan an LM Studio models directory for model files.
|
|
|
|
LM Studio uses a ``publisher/model-name`` folder structure containing
|
|
GGUF files, or standalone GGUF files at the top level.
|
|
"""
|
|
if not lm_dir.exists() or not lm_dir.is_dir():
|
|
return []
|
|
|
|
found: List[LocalModelInfo] = []
|
|
for child in lm_dir.iterdir():
|
|
if not child.is_dir():
|
|
if child.suffix == ".gguf" and child.is_file():
|
|
try:
|
|
updated_at = child.stat().st_mtime
|
|
except OSError:
|
|
updated_at = None
|
|
found.append(
|
|
LocalModelInfo(
|
|
id = str(child),
|
|
display_name = child.stem,
|
|
path = str(child),
|
|
source = "lmstudio",
|
|
updated_at = updated_at,
|
|
),
|
|
)
|
|
continue
|
|
|
|
# child is a publisher directory — scan its sub-directories
|
|
for model_dir in child.iterdir():
|
|
if model_dir.is_dir():
|
|
has_model = (
|
|
any(model_dir.glob("*.gguf"))
|
|
or (model_dir / "config.json").exists()
|
|
or any(model_dir.glob("*.safetensors"))
|
|
)
|
|
if not has_model:
|
|
continue
|
|
model_id = f"{child.name}/{model_dir.name}"
|
|
try:
|
|
updated_at = model_dir.stat().st_mtime
|
|
except OSError:
|
|
updated_at = None
|
|
found.append(
|
|
LocalModelInfo(
|
|
id = str(model_dir),
|
|
model_id = model_id,
|
|
display_name = model_dir.name,
|
|
path = str(model_dir),
|
|
source = "lmstudio",
|
|
updated_at = updated_at,
|
|
),
|
|
)
|
|
elif model_dir.suffix == ".gguf" and model_dir.is_file():
|
|
try:
|
|
updated_at = model_dir.stat().st_mtime
|
|
except OSError:
|
|
updated_at = None
|
|
found.append(
|
|
LocalModelInfo(
|
|
id = str(model_dir),
|
|
display_name = model_dir.stem,
|
|
path = str(model_dir),
|
|
source = "lmstudio",
|
|
updated_at = updated_at,
|
|
),
|
|
)
|
|
return found
|
|
|
|
|
|
@router.get("/local", response_model = LocalModelListResponse)
|
|
async def list_local_models(
|
|
models_dir: str = Query(
|
|
default = "./models", description = "Directory to scan for local model folders"
|
|
),
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""
|
|
List local model candidates from custom models dir, HF cache,
|
|
legacy Unsloth HF cache, and LM Studio directories.
|
|
"""
|
|
from utils.paths import (
|
|
legacy_hf_cache_dir,
|
|
hf_default_cache_dir,
|
|
lmstudio_model_dirs,
|
|
)
|
|
|
|
# Resolve all scan directories up front.
|
|
hf_cache_dir = _resolve_hf_cache_dir()
|
|
legacy_hf = legacy_hf_cache_dir()
|
|
hf_default = hf_default_cache_dir()
|
|
lm_dirs = lmstudio_model_dirs()
|
|
|
|
# Validate models_dir against an allowlist of trusted directories.
|
|
# Only the trusted Path objects are used for filesystem access -- the
|
|
# user-supplied string is only used for matching, never for path construction.
|
|
allowed_roots: list[Path] = [Path("./models").resolve(), hf_cache_dir]
|
|
if legacy_hf.is_dir():
|
|
allowed_roots.append(legacy_hf)
|
|
if hf_default.is_dir():
|
|
allowed_roots.append(hf_default)
|
|
try:
|
|
from utils.paths import studio_root, outputs_root
|
|
|
|
allowed_roots.extend([studio_root(), outputs_root()])
|
|
except Exception:
|
|
pass
|
|
|
|
requested = os.path.realpath(os.path.expanduser(models_dir))
|
|
models_root = None
|
|
for root in allowed_roots:
|
|
root_str = os.path.realpath(str(root))
|
|
if requested == root_str or requested.startswith(root_str + os.sep):
|
|
models_root = root # Use the trusted root, not the user-supplied path
|
|
break
|
|
if models_root is None:
|
|
raise HTTPException(
|
|
status_code = 403,
|
|
detail = "Directory not allowed",
|
|
)
|
|
|
|
try:
|
|
local_models = _scan_models_dir(models_root) + _scan_hf_cache(hf_cache_dir)
|
|
|
|
# Scan legacy Unsloth HF cache for backward compatibility
|
|
if legacy_hf.is_dir() and legacy_hf.resolve() != hf_cache_dir.resolve():
|
|
local_models += _scan_hf_cache(legacy_hf)
|
|
|
|
# Scan HF system default cache (may differ when env vars are overridden)
|
|
if (
|
|
hf_default.is_dir()
|
|
and hf_default.resolve() != hf_cache_dir.resolve()
|
|
and hf_default.resolve() != legacy_hf.resolve()
|
|
):
|
|
local_models += _scan_hf_cache(hf_default)
|
|
|
|
# Scan LM Studio directories
|
|
for lm_dir in lm_dirs:
|
|
local_models += _scan_lmstudio_dir(lm_dir)
|
|
|
|
deduped: dict[str, LocalModelInfo] = {}
|
|
for model in local_models:
|
|
if model.id not in deduped:
|
|
deduped[model.id] = model
|
|
|
|
models = sorted(
|
|
deduped.values(),
|
|
key = lambda item: (item.updated_at or 0),
|
|
reverse = True,
|
|
)
|
|
|
|
return LocalModelListResponse(
|
|
models_dir = str(models_root),
|
|
hf_cache_dir = str(hf_cache_dir),
|
|
lmstudio_dirs = [str(d) for d in lm_dirs],
|
|
models = models,
|
|
)
|
|
except Exception as e:
|
|
logger.error(f"Error listing local models: {e}", exc_info = True)
|
|
raise HTTPException(
|
|
status_code = 500,
|
|
detail = f"Failed to list local models: {str(e)}",
|
|
)
|
|
|
|
|
|
@router.get("/list")
|
|
async def list_models(
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""
|
|
List available models (default models and loaded models).
|
|
|
|
This endpoint returns the default models and any currently loaded models.
|
|
"""
|
|
try:
|
|
inference_backend = get_inference_backend()
|
|
|
|
# Get default models
|
|
default_models = inference_backend.default_models
|
|
|
|
# Get loaded models
|
|
loaded_models = []
|
|
for model_name, model_data in inference_backend.models.items():
|
|
_is_vision = model_data.get("is_vision", False)
|
|
_audio_type = model_data.get("audio_type")
|
|
model_info = ModelDetails(
|
|
id = model_name,
|
|
name = model_name.split("/")[-1] if "/" in model_name else model_name,
|
|
is_vision = _is_vision,
|
|
is_lora = model_data.get("is_lora", False),
|
|
is_audio = model_data.get("is_audio", False),
|
|
audio_type = _audio_type,
|
|
has_audio_input = model_data.get("has_audio_input", False),
|
|
model_type = derive_model_type(_is_vision, _audio_type),
|
|
)
|
|
loaded_models.append(model_info)
|
|
|
|
# Include active GGUF model (loaded via llama-server)
|
|
from routes.inference import get_llama_cpp_backend
|
|
|
|
llama_backend = get_llama_cpp_backend()
|
|
if llama_backend.is_loaded and llama_backend.model_identifier:
|
|
loaded_models.append(
|
|
ModelDetails(
|
|
id = llama_backend.model_identifier,
|
|
name = llama_backend.model_identifier.split("/")[-1],
|
|
is_gguf = True,
|
|
is_vision = llama_backend.is_vision,
|
|
is_audio = getattr(llama_backend, "_is_audio", False),
|
|
audio_type = getattr(llama_backend, "_audio_type", None),
|
|
)
|
|
)
|
|
|
|
# Combine default and loaded models
|
|
all_models = []
|
|
seen_ids = set()
|
|
|
|
# Add default models
|
|
for model_id in default_models:
|
|
if model_id not in seen_ids:
|
|
model_info = ModelDetails(
|
|
id = model_id,
|
|
name = model_id.split("/")[-1] if "/" in model_id else model_id,
|
|
is_gguf = model_id.upper().endswith("-GGUF"),
|
|
)
|
|
all_models.append(model_info)
|
|
seen_ids.add(model_id)
|
|
|
|
# Add loaded models
|
|
for model_info in loaded_models:
|
|
if model_info.id not in seen_ids:
|
|
all_models.append(model_info)
|
|
seen_ids.add(model_info.id)
|
|
|
|
return ModelListResponse(models = all_models, default_models = default_models)
|
|
|
|
except Exception as e:
|
|
logger.error(f"Error listing models: {e}", exc_info = True)
|
|
raise HTTPException(status_code = 500, detail = f"Failed to list models: {str(e)}")
|
|
|
|
|
|
def _get_max_position_embeddings(config) -> Optional[int]:
|
|
"""Extract max_position_embeddings from a model config, checking text_config fallback."""
|
|
if hasattr(config, "max_position_embeddings"):
|
|
return config.max_position_embeddings
|
|
if hasattr(config, "text_config") and hasattr(
|
|
config.text_config, "max_position_embeddings"
|
|
):
|
|
return config.text_config.max_position_embeddings
|
|
return None
|
|
|
|
|
|
def _get_model_size_bytes(
|
|
model_name: str, hf_token: Optional[str] = None
|
|
) -> Optional[int]:
|
|
"""Get total size of model weight files from HF Hub."""
|
|
try:
|
|
from huggingface_hub import HfApi
|
|
|
|
api = HfApi(token = hf_token)
|
|
info = api.repo_info(model_name, repo_type = "model", token = hf_token)
|
|
if not info.siblings:
|
|
return None
|
|
|
|
weight_exts = (".safetensors", ".bin", ".pt", ".pth", ".gguf")
|
|
total = 0
|
|
for sibling in info.siblings:
|
|
if sibling.rfilename and any(
|
|
sibling.rfilename.endswith(ext) for ext in weight_exts
|
|
):
|
|
if sibling.size is not None:
|
|
total += sibling.size
|
|
|
|
return total if total > 0 else None
|
|
except Exception as e:
|
|
logger.warning(f"Could not get model size for {model_name}: {e}")
|
|
return None
|
|
|
|
|
|
@router.get("/config/{model_name:path}")
|
|
async def get_model_config(
|
|
model_name: str,
|
|
hf_token: Optional[str] = Query(None),
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""
|
|
Get configuration for a specific model.
|
|
|
|
This endpoint wraps the backend load_model_defaults function.
|
|
"""
|
|
try:
|
|
from utils.models.model_config import is_local_path
|
|
|
|
if not is_local_path(model_name):
|
|
model_name = model_name.lower()
|
|
|
|
logger.info(f"Getting model config for: {model_name}")
|
|
from utils.models.model_config import detect_audio_type
|
|
|
|
# Load model defaults from backend
|
|
config_dict = load_model_defaults(model_name)
|
|
|
|
# Detect model capabilities (pass HF token for gated models)
|
|
is_vision = is_vision_model(model_name)
|
|
is_embedding = is_embedding_model(model_name, hf_token = hf_token)
|
|
audio_type = detect_audio_type(model_name, hf_token = hf_token)
|
|
|
|
# Check if it's a LoRA adapter
|
|
is_lora = False
|
|
base_model = None
|
|
max_position_embeddings = None
|
|
try:
|
|
model_config = ModelConfig.from_identifier(model_name)
|
|
is_lora = model_config.is_lora
|
|
base_model = model_config.base_model if is_lora else None
|
|
max_position_embeddings = _get_max_position_embeddings(model_config)
|
|
except Exception:
|
|
pass
|
|
|
|
# Fallback: try AutoConfig directly if not found yet
|
|
if max_position_embeddings is None:
|
|
try:
|
|
from transformers import AutoConfig as _AutoConfig
|
|
|
|
_trust = model_name.lower().startswith("unsloth/")
|
|
_ac = _AutoConfig.from_pretrained(
|
|
model_name, trust_remote_code = _trust, token = hf_token
|
|
)
|
|
max_position_embeddings = _get_max_position_embeddings(_ac)
|
|
except Exception:
|
|
pass
|
|
|
|
logger.info(
|
|
f"Model config result for {model_name}: is_vision={is_vision}, is_embedding={is_embedding}, audio_type={audio_type}, is_lora={is_lora}, max_position_embeddings={max_position_embeddings}"
|
|
)
|
|
return ModelDetails(
|
|
id = model_name,
|
|
model_name = model_name,
|
|
config = config_dict,
|
|
is_vision = is_vision,
|
|
is_embedding = is_embedding,
|
|
is_lora = is_lora,
|
|
is_audio = audio_type is not None,
|
|
audio_type = audio_type,
|
|
has_audio_input = is_audio_input_type(audio_type),
|
|
model_type = derive_model_type(is_vision, audio_type, is_embedding),
|
|
base_model = base_model,
|
|
max_position_embeddings = max_position_embeddings,
|
|
model_size_bytes = _get_model_size_bytes(model_name, hf_token),
|
|
)
|
|
|
|
except Exception as e:
|
|
logger.error(f"Error getting model config: {e}", exc_info = True)
|
|
raise HTTPException(
|
|
status_code = 500, detail = f"Failed to get model config: {str(e)}"
|
|
)
|
|
|
|
|
|
@router.get("/loras")
|
|
async def scan_loras(
|
|
outputs_dir: str = Query(
|
|
default = str(outputs_root()), description = "Directory to scan for LoRA adapters"
|
|
),
|
|
exports_dir: str = Query(
|
|
default = str(exports_root()), description = "Directory to scan for exported models"
|
|
),
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""
|
|
Scan for trained LoRA adapters and exported models.
|
|
|
|
Returns both training outputs (from outputs_dir) and exported models
|
|
(from exports_dir) in a single list, distinguished by source field.
|
|
"""
|
|
try:
|
|
resolved_outputs_dir = str(resolve_output_dir(outputs_dir))
|
|
resolved_exports_dir = str(resolve_export_dir(exports_dir))
|
|
lora_list = []
|
|
|
|
# Scan training outputs
|
|
trained_loras = scan_trained_loras(outputs_dir = resolved_outputs_dir)
|
|
for display_name, adapter_path in trained_loras:
|
|
base_model = get_base_model_from_lora(adapter_path)
|
|
lora_list.append(
|
|
LoRAInfo(
|
|
display_name = display_name,
|
|
adapter_path = adapter_path,
|
|
base_model = base_model,
|
|
source = "training",
|
|
)
|
|
)
|
|
|
|
# Scan exported models (merged, LoRA, base — skips GGUF)
|
|
exported = scan_exported_models(exports_dir = resolved_exports_dir)
|
|
for display_name, model_path, export_type, base_model in exported:
|
|
lora_list.append(
|
|
LoRAInfo(
|
|
display_name = display_name,
|
|
adapter_path = model_path,
|
|
base_model = base_model,
|
|
source = "exported",
|
|
export_type = export_type,
|
|
)
|
|
)
|
|
|
|
return LoRAScanResponse(loras = lora_list, outputs_dir = resolved_outputs_dir)
|
|
|
|
except Exception as e:
|
|
logger.error(f"Error scanning LoRAs: {e}", exc_info = True)
|
|
raise HTTPException(
|
|
status_code = 500, detail = f"Failed to scan LoRA adapters: {str(e)}"
|
|
)
|
|
|
|
|
|
@router.get("/loras/{lora_path:path}/base-model", response_model = LoRABaseModelResponse)
|
|
async def get_lora_base_model(
|
|
lora_path: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""
|
|
Get the base model for a LoRA adapter.
|
|
|
|
This endpoint wraps the backend get_base_model_from_lora function.
|
|
"""
|
|
try:
|
|
base_model = get_base_model_from_lora(lora_path)
|
|
|
|
if base_model is None:
|
|
raise HTTPException(
|
|
status_code = 404,
|
|
detail = f"Could not determine base model for LoRA: {lora_path}",
|
|
)
|
|
|
|
return LoRABaseModelResponse(
|
|
lora_path = lora_path,
|
|
base_model = base_model,
|
|
)
|
|
|
|
except HTTPException:
|
|
raise
|
|
except Exception as e:
|
|
logger.error(f"Error getting LoRA base model: {e}", exc_info = True)
|
|
raise HTTPException(
|
|
status_code = 500, detail = f"Failed to get base model: {str(e)}"
|
|
)
|
|
|
|
|
|
@router.get("/check-vision/{model_name:path}", response_model = VisionCheckResponse)
|
|
async def check_vision_model(
|
|
model_name: str,
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""
|
|
Check if a model is a vision model.
|
|
|
|
This endpoint wraps the backend is_vision_model function.
|
|
"""
|
|
try:
|
|
logger.info(f"Checking if vision model: {model_name}")
|
|
is_vision = is_vision_model(model_name)
|
|
|
|
logger.info(f"Vision check result for {model_name}: is_vision={is_vision}")
|
|
return VisionCheckResponse(
|
|
model_name = model_name,
|
|
is_vision = is_vision,
|
|
)
|
|
|
|
except Exception as e:
|
|
logger.error(f"Error checking vision model: {e}", exc_info = True)
|
|
raise HTTPException(
|
|
status_code = 500, detail = f"Failed to check vision model: {str(e)}"
|
|
)
|
|
|
|
|
|
@router.get("/check-embedding/{model_name:path}", response_model = EmbeddingCheckResponse)
|
|
async def check_embedding_model(
|
|
model_name: str,
|
|
hf_token: Optional[str] = Query(None),
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""
|
|
Check if a model is an embedding model.
|
|
|
|
This endpoint wraps the backend is_embedding_model function.
|
|
"""
|
|
try:
|
|
logger.info(f"Checking if embedding model: {model_name}")
|
|
is_embedding = is_embedding_model(model_name, hf_token = hf_token)
|
|
|
|
logger.info(
|
|
f"Embedding check result for {model_name}: is_embedding={is_embedding}"
|
|
)
|
|
return EmbeddingCheckResponse(
|
|
model_name = model_name,
|
|
is_embedding = is_embedding,
|
|
)
|
|
|
|
except Exception as e:
|
|
logger.error(f"Error checking embedding model: {e}", exc_info = True)
|
|
raise HTTPException(
|
|
status_code = 500, detail = f"Failed to check embedding model: {str(e)}"
|
|
)
|
|
|
|
|
|
@router.get("/gguf-variants", response_model = GgufVariantsResponse)
|
|
async def get_gguf_variants(
|
|
repo_id: str = Query(
|
|
..., description = "HuggingFace repo ID (e.g. 'unsloth/gemma-3-4b-it-GGUF')"
|
|
),
|
|
hf_token: Optional[str] = Query(
|
|
None, description = "HuggingFace token for private repos"
|
|
),
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""
|
|
List available GGUF quantization variants for a HuggingFace repo.
|
|
|
|
Returns all available quantization variants (Q4_K_M, Q8_0, BF16, etc.)
|
|
with file sizes, whether the model supports vision, and the recommended
|
|
default variant.
|
|
"""
|
|
try:
|
|
variants, has_vision = list_gguf_variants(repo_id, hf_token = hf_token)
|
|
|
|
# Determine default variant
|
|
filenames = [v.filename for v in variants]
|
|
best = _pick_best_gguf(filenames)
|
|
default_variant = _extract_quant_label(best) if best else None
|
|
|
|
# Check which variants are fully downloaded in the HF cache.
|
|
# For split GGUFs, ALL shards must be present -- sum cached bytes
|
|
# per variant and compare against the expected total.
|
|
# HF cache dir uses the exact case from the repo_id at download time,
|
|
# which may differ from the canonical HF repo_id, so do a
|
|
# case-insensitive match.
|
|
cached_bytes_by_quant: dict[str, int] = {}
|
|
try:
|
|
import re as _re
|
|
from huggingface_hub import constants as hf_constants
|
|
|
|
# Sanitize repo_id: must be "owner/name" with safe chars only
|
|
if not _is_valid_repo_id(repo_id):
|
|
raise ValueError(f"Invalid repo_id format: {repo_id}")
|
|
|
|
cache_dir = Path(hf_constants.HF_HUB_CACHE)
|
|
target = f"models--{repo_id.replace('/', '--')}".lower()
|
|
for entry in cache_dir.iterdir():
|
|
if entry.name.lower() == target:
|
|
snapshots = entry / "snapshots"
|
|
if snapshots.is_dir():
|
|
for snap in snapshots.iterdir():
|
|
for f in snap.rglob("*.gguf"):
|
|
q = _extract_quant_label(f.name)
|
|
cached_bytes_by_quant[q] = (
|
|
cached_bytes_by_quant.get(q, 0) + f.stat().st_size
|
|
)
|
|
break
|
|
except Exception:
|
|
pass
|
|
|
|
def _is_fully_downloaded(variant) -> bool:
|
|
cached = cached_bytes_by_quant.get(variant.quant, 0)
|
|
if cached == 0 or variant.size_bytes == 0:
|
|
return False
|
|
# Allow small rounding tolerance (symlinks vs real sizes)
|
|
return cached >= variant.size_bytes * 0.99
|
|
|
|
return GgufVariantsResponse(
|
|
repo_id = repo_id,
|
|
variants = [
|
|
GgufVariantDetail(
|
|
filename = v.filename,
|
|
quant = v.quant,
|
|
size_bytes = v.size_bytes,
|
|
downloaded = _is_fully_downloaded(v),
|
|
)
|
|
for v in variants
|
|
],
|
|
has_vision = has_vision,
|
|
default_variant = default_variant,
|
|
)
|
|
|
|
except Exception as e:
|
|
logger.error(f"Error listing GGUF variants for '{repo_id}': {e}", exc_info = True)
|
|
raise HTTPException(
|
|
status_code = 500,
|
|
detail = f"Failed to list GGUF variants: {str(e)}",
|
|
)
|
|
|
|
|
|
@router.get("/gguf-download-progress")
|
|
async def get_gguf_download_progress(
|
|
repo_id: str = Query(..., description = "HuggingFace repo ID"),
|
|
variant: str = Query("", description = "Quantization variant (e.g. UD-TQ1_0)"),
|
|
expected_bytes: int = Query(0, description = "Expected total download size in bytes"),
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""Return download progress by checking cached GGUF files for a specific variant.
|
|
|
|
Tracks completed shard downloads in snapshots and in-progress downloads
|
|
in the blobs directory (incomplete files).
|
|
"""
|
|
try:
|
|
if not _is_valid_repo_id(repo_id):
|
|
return {
|
|
"downloaded_bytes": 0,
|
|
"expected_bytes": expected_bytes,
|
|
"progress": 0,
|
|
}
|
|
|
|
from huggingface_hub import constants as hf_constants
|
|
|
|
cache_dir = Path(hf_constants.HF_HUB_CACHE)
|
|
target = f"models--{repo_id.replace('/', '--')}".lower()
|
|
variant_lower = variant.lower().replace("-", "").replace("_", "")
|
|
downloaded_bytes = 0
|
|
in_progress_bytes = 0
|
|
for entry in cache_dir.iterdir():
|
|
if entry.name.lower() == target:
|
|
# Count completed .gguf files matching this variant in snapshots
|
|
for f in entry.rglob("*.gguf"):
|
|
fname = f.name.lower().replace("-", "").replace("_", "")
|
|
if not variant_lower or variant_lower in fname:
|
|
downloaded_bytes += f.stat().st_size
|
|
# Check blobs for in-progress downloads (.incomplete files)
|
|
blobs_dir = entry / "blobs"
|
|
if blobs_dir.is_dir():
|
|
for f in blobs_dir.iterdir():
|
|
if f.is_file() and f.name.endswith(".incomplete"):
|
|
in_progress_bytes += f.stat().st_size
|
|
break
|
|
|
|
total_progress_bytes = downloaded_bytes + in_progress_bytes
|
|
progress = (
|
|
min(total_progress_bytes / expected_bytes, 0.99)
|
|
if expected_bytes > 0
|
|
else 0
|
|
)
|
|
# Only report 1.0 when all bytes are in completed files (not in-progress)
|
|
if expected_bytes > 0 and downloaded_bytes >= expected_bytes:
|
|
progress = 1.0
|
|
return {
|
|
"downloaded_bytes": total_progress_bytes,
|
|
"expected_bytes": expected_bytes,
|
|
"progress": round(progress, 3),
|
|
}
|
|
except Exception:
|
|
return {"downloaded_bytes": 0, "expected_bytes": expected_bytes, "progress": 0}
|
|
|
|
|
|
@router.get("/download-progress")
|
|
async def get_download_progress(
|
|
repo_id: str = Query(..., description = "HuggingFace repo ID"),
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""Return download progress for any HuggingFace model repo.
|
|
|
|
Checks the local HF cache for completed blobs and in-progress
|
|
(.incomplete) downloads. Uses the HF API to determine the expected
|
|
total size on the first call, then caches it for subsequent polls.
|
|
"""
|
|
_empty = {"downloaded_bytes": 0, "expected_bytes": 0, "progress": 0}
|
|
try:
|
|
if not _is_valid_repo_id(repo_id):
|
|
return _empty
|
|
|
|
from huggingface_hub import constants as hf_constants
|
|
|
|
cache_dir = Path(hf_constants.HF_HUB_CACHE)
|
|
target = f"models--{repo_id.replace('/', '--')}".lower()
|
|
completed_bytes = 0
|
|
in_progress_bytes = 0
|
|
|
|
for entry in cache_dir.iterdir():
|
|
if entry.name.lower() != target:
|
|
continue
|
|
blobs_dir = entry / "blobs"
|
|
if not blobs_dir.is_dir():
|
|
break
|
|
for f in blobs_dir.iterdir():
|
|
if not f.is_file():
|
|
continue
|
|
if f.name.endswith(".incomplete"):
|
|
in_progress_bytes += f.stat().st_size
|
|
else:
|
|
completed_bytes += f.stat().st_size
|
|
break
|
|
|
|
downloaded_bytes = completed_bytes + in_progress_bytes
|
|
if downloaded_bytes == 0:
|
|
return _empty
|
|
|
|
# Get expected size from HF API (cached per repo_id)
|
|
expected_bytes = _get_repo_size_cached(repo_id)
|
|
if expected_bytes <= 0:
|
|
# Cannot determine total; report bytes only, no percentage
|
|
return {
|
|
"downloaded_bytes": downloaded_bytes,
|
|
"expected_bytes": 0,
|
|
"progress": 0,
|
|
}
|
|
|
|
# Use 95% threshold for completion (blob deduplication can make
|
|
# completed_bytes differ slightly from expected_bytes).
|
|
# Do NOT use "no .incomplete files" as a completion signal --
|
|
# HF downloads files sequentially, so between files there are
|
|
# no .incomplete files even though the download is far from done.
|
|
if completed_bytes >= expected_bytes * 0.95:
|
|
progress = 1.0
|
|
else:
|
|
progress = min(downloaded_bytes / expected_bytes, 0.99)
|
|
return {
|
|
"downloaded_bytes": downloaded_bytes,
|
|
"expected_bytes": expected_bytes,
|
|
"progress": round(progress, 3),
|
|
}
|
|
except Exception as e:
|
|
logger.warning(f"Error checking download progress for {repo_id}: {e}")
|
|
return _empty
|
|
|
|
|
|
_repo_size_cache: dict[str, int] = {}
|
|
|
|
|
|
def _get_repo_size_cached(repo_id: str) -> int:
|
|
if repo_id in _repo_size_cache:
|
|
return _repo_size_cache[repo_id]
|
|
try:
|
|
from huggingface_hub import model_info as hf_model_info
|
|
|
|
info = hf_model_info(repo_id, token = None, files_metadata = True)
|
|
total = sum(s.size for s in info.siblings if s.size)
|
|
_repo_size_cache[repo_id] = total
|
|
return total
|
|
except Exception as e:
|
|
logger.warning(f"Failed to get repo size for {repo_id}: {e}")
|
|
return 0
|
|
|
|
|
|
def _all_hf_cache_scans():
|
|
"""Return scan_cache_dir results for the active, legacy, and default HF caches."""
|
|
from huggingface_hub import scan_cache_dir
|
|
from utils.paths import legacy_hf_cache_dir, hf_default_cache_dir
|
|
|
|
scans = [scan_cache_dir()]
|
|
seen: set[str] = set()
|
|
try:
|
|
# Resolve the active cache dir so we can dedup
|
|
from huggingface_hub.constants import HF_HUB_CACHE
|
|
|
|
seen.add(str(Path(HF_HUB_CACHE).resolve()))
|
|
except Exception:
|
|
pass
|
|
|
|
for extra_fn in (legacy_hf_cache_dir, hf_default_cache_dir):
|
|
extra = extra_fn()
|
|
if extra.is_dir() and str(extra.resolve()) not in seen:
|
|
seen.add(str(extra.resolve()))
|
|
try:
|
|
scans.append(scan_cache_dir(cache_dir = str(extra)))
|
|
except Exception as exc:
|
|
logger.warning("Could not scan HF cache %s: %s", extra, exc)
|
|
return scans
|
|
|
|
|
|
@router.get("/cached-gguf")
|
|
async def list_cached_gguf(
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""List GGUF repos downloaded to HF cache, legacy Unsloth cache, and HF default cache."""
|
|
try:
|
|
cache_scans = _all_hf_cache_scans()
|
|
|
|
seen_lower: dict[str, dict] = {}
|
|
for hf_cache in cache_scans:
|
|
for repo_info in hf_cache.repos:
|
|
if repo_info.repo_type != "model":
|
|
continue
|
|
repo_id = repo_info.repo_id
|
|
if not repo_id.upper().endswith("-GGUF"):
|
|
continue
|
|
total_size = 0
|
|
has_gguf = False
|
|
for revision in repo_info.revisions:
|
|
for f in revision.files:
|
|
if f.file_name.endswith(".gguf"):
|
|
has_gguf = True
|
|
total_size += f.size_on_disk
|
|
if not has_gguf:
|
|
continue
|
|
key = repo_id.lower()
|
|
existing = seen_lower.get(key)
|
|
if existing is None or total_size > existing["size_bytes"]:
|
|
seen_lower[key] = {
|
|
"repo_id": repo_id,
|
|
"size_bytes": total_size,
|
|
"cache_path": str(repo_info.repo_path),
|
|
}
|
|
cached = sorted(seen_lower.values(), key = lambda c: c["repo_id"])
|
|
return {"cached": cached}
|
|
except Exception as e:
|
|
logger.error(f"Error listing cached GGUF repos: {e}", exc_info = True)
|
|
return {"cached": []}
|
|
|
|
|
|
@router.get("/cached-models")
|
|
async def list_cached_models(
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""List non-GGUF model repos downloaded to HF cache, legacy Unsloth cache, and HF default cache."""
|
|
_WEIGHT_EXTENSIONS = (".safetensors", ".bin")
|
|
|
|
try:
|
|
cache_scans = _all_hf_cache_scans()
|
|
|
|
seen_lower: dict[str, dict] = {}
|
|
for hf_cache in cache_scans:
|
|
for repo_info in hf_cache.repos:
|
|
if repo_info.repo_type != "model":
|
|
continue
|
|
repo_id = repo_info.repo_id
|
|
if repo_id.upper().endswith("-GGUF"):
|
|
continue
|
|
total_size = sum(
|
|
f.size_on_disk for rev in repo_info.revisions for f in rev.files
|
|
)
|
|
if total_size == 0:
|
|
continue
|
|
has_weights = any(
|
|
f.file_name.endswith(_WEIGHT_EXTENSIONS)
|
|
for rev in repo_info.revisions
|
|
for f in rev.files
|
|
)
|
|
if not has_weights:
|
|
continue
|
|
key = repo_id.lower()
|
|
existing = seen_lower.get(key)
|
|
if existing is None or total_size > existing["size_bytes"]:
|
|
seen_lower[key] = {
|
|
"repo_id": repo_id,
|
|
"size_bytes": total_size,
|
|
}
|
|
cached = sorted(seen_lower.values(), key = lambda c: c["repo_id"])
|
|
return {"cached": cached}
|
|
except Exception as e:
|
|
logger.error(f"Error listing cached models: {e}", exc_info = True)
|
|
return {"cached": []}
|
|
|
|
|
|
@router.delete("/delete-cached")
|
|
async def delete_cached_model(
|
|
repo_id: str = Body(...),
|
|
variant: Optional[str] = Body(None),
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""Delete a cached model repo (or a specific GGUF variant) from the HF cache.
|
|
|
|
When *variant* is provided, only the GGUF files matching that quant label
|
|
are removed (e.g. ``UD-Q4_K_XL``). Otherwise the entire repo is deleted.
|
|
Refuses if the model is currently loaded for inference.
|
|
"""
|
|
if not _is_valid_repo_id(repo_id):
|
|
raise HTTPException(status_code = 400, detail = "Invalid repo_id format")
|
|
|
|
# Check if model is currently loaded
|
|
try:
|
|
from routes.inference import get_llama_cpp_backend
|
|
|
|
llama_backend = get_llama_cpp_backend()
|
|
if llama_backend.is_loaded and llama_backend.model_identifier:
|
|
loaded_id = llama_backend.model_identifier.lower()
|
|
if loaded_id == repo_id.lower() or loaded_id.startswith(repo_id.lower()):
|
|
raise HTTPException(
|
|
status_code = 400,
|
|
detail = "Unload the model before deleting",
|
|
)
|
|
except HTTPException:
|
|
raise
|
|
except Exception:
|
|
pass
|
|
|
|
try:
|
|
inference_backend = get_inference_backend()
|
|
if inference_backend.active_model_name:
|
|
active = inference_backend.active_model_name.lower()
|
|
if active == repo_id.lower() or active.startswith(repo_id.lower()):
|
|
raise HTTPException(
|
|
status_code = 400,
|
|
detail = "Unload the model before deleting",
|
|
)
|
|
except HTTPException:
|
|
raise
|
|
except Exception:
|
|
pass
|
|
|
|
try:
|
|
cache_scans = _all_hf_cache_scans()
|
|
|
|
target_repo = None
|
|
for hf_cache in cache_scans:
|
|
for repo_info in hf_cache.repos:
|
|
if repo_info.repo_type != "model":
|
|
continue
|
|
if repo_info.repo_id.lower() == repo_id.lower():
|
|
target_repo = repo_info
|
|
break
|
|
if target_repo is not None:
|
|
break
|
|
|
|
if target_repo is None:
|
|
raise HTTPException(status_code = 404, detail = "Model not found in cache")
|
|
|
|
# ── Per-variant GGUF deletion ────────────────────────────
|
|
if variant:
|
|
deleted_bytes = 0
|
|
deleted_count = 0
|
|
for rev in target_repo.revisions:
|
|
for f in rev.files:
|
|
if not f.file_name.endswith(".gguf"):
|
|
continue
|
|
quant = _extract_quant_label(f.file_name)
|
|
if quant.lower() != variant.lower():
|
|
continue
|
|
# Delete the blob (actual data) and the snapshot symlink
|
|
try:
|
|
blob = Path(f.blob_path)
|
|
snap = Path(f.file_path)
|
|
size = blob.stat().st_size if blob.exists() else 0
|
|
if snap.exists() or snap.is_symlink():
|
|
snap.unlink()
|
|
if blob.exists():
|
|
blob.unlink()
|
|
deleted_bytes += size
|
|
deleted_count += 1
|
|
except Exception as e:
|
|
logger.warning(f"Failed to delete {f.file_name}: {e}")
|
|
|
|
if deleted_count == 0:
|
|
raise HTTPException(
|
|
status_code = 404,
|
|
detail = f"Variant {variant} not found in cache for {repo_id}",
|
|
)
|
|
|
|
freed_mb = deleted_bytes / (1024 * 1024)
|
|
logger.info(
|
|
f"Deleted {deleted_count} file(s) for {repo_id} variant {variant}: "
|
|
f"{freed_mb:.1f} MB freed"
|
|
)
|
|
return {"status": "deleted", "repo_id": repo_id, "variant": variant}
|
|
|
|
# ── Full repo deletion ───────────────────────────────────
|
|
revision_hashes = [rev.commit_hash for rev in target_repo.revisions]
|
|
if not revision_hashes:
|
|
raise HTTPException(status_code = 404, detail = "No revisions found for model")
|
|
|
|
delete_strategy = hf_cache.delete_revisions(*revision_hashes)
|
|
logger.info(
|
|
f"Deleting cached model {repo_id}: "
|
|
f"{delete_strategy.expected_freed_size_str} will be freed"
|
|
)
|
|
delete_strategy.execute()
|
|
|
|
return {"status": "deleted", "repo_id": repo_id}
|
|
|
|
except HTTPException:
|
|
raise
|
|
except Exception as e:
|
|
logger.error(f"Error deleting cached model {repo_id}: {e}", exc_info = True)
|
|
raise HTTPException(
|
|
status_code = 500,
|
|
detail = f"Failed to delete cached model: {str(e)}",
|
|
)
|
|
|
|
|
|
@router.get("/checkpoints", response_model = CheckpointListResponse)
|
|
async def list_checkpoints(
|
|
outputs_dir: str = Query(
|
|
default = str(outputs_root()),
|
|
description = "Directory to scan for checkpoints",
|
|
),
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""
|
|
List available checkpoints in the outputs directory.
|
|
|
|
Scans the outputs folder for training runs and their checkpoints.
|
|
"""
|
|
try:
|
|
resolved_outputs_dir = str(resolve_output_dir(outputs_dir))
|
|
raw_models = scan_checkpoints(outputs_dir = resolved_outputs_dir)
|
|
|
|
models = [
|
|
ModelCheckpoints(
|
|
name = model_name,
|
|
checkpoints = [
|
|
CheckpointInfo(display_name = display_name, path = path, loss = loss)
|
|
for display_name, path, loss in checkpoints
|
|
],
|
|
base_model = metadata.get("base_model"),
|
|
peft_type = metadata.get("peft_type"),
|
|
lora_rank = metadata.get("lora_rank"),
|
|
is_quantized = metadata.get("is_quantized", False),
|
|
)
|
|
for model_name, checkpoints, metadata in raw_models
|
|
]
|
|
|
|
return CheckpointListResponse(
|
|
outputs_dir = resolved_outputs_dir,
|
|
models = models,
|
|
)
|
|
except Exception as e:
|
|
logger.error(f"Error listing checkpoints: {e}", exc_info = True)
|
|
raise HTTPException(
|
|
status_code = 500,
|
|
detail = f"Failed to list checkpoints: {str(e)}",
|
|
)
|