gaussian_studio / utils /hf_utils.py
dgarch424's picture
Upload 21 files
728fc83 verified
Raw
History Blame Contribute Delete
3.78 kB
"""
utils/hf_utils.py
─────────────────
Utilities for validating and resolving HuggingFace model IDs or external URLs.
Used by the UI's "Custom model" text input before attempting to load.
"""
from __future__ import annotations
import logging
import re
from typing import Literal
import requests
from huggingface_hub import HfApi, model_info
from huggingface_hub.utils import RepositoryNotFoundError, HFValidationError
logger = logging.getLogger(__name__)
HF_REPO_RE = re.compile(r"^[A-Za-z0-9_\-\.]+/[A-Za-z0-9_\-\.]+$")
URL_RE = re.compile(r"^https?://")
ModelSource = Literal["hf_repo", "url", "invalid"]
def classify_model_id(model_id: str) -> ModelSource:
"""
Decide whether a raw string is:
- a valid HuggingFace repo slug ("org/model-name")
- an external URL ("https://…")
- invalid
"""
model_id = model_id.strip()
if URL_RE.match(model_id):
return "url"
if HF_REPO_RE.match(model_id):
return "hf_repo"
return "invalid"
def validate_hf_repo(repo_id: str, token: str | None = None) -> tuple[bool, str]:
"""
Check whether a HuggingFace repo exists and is publicly accessible.
Returns
-------
(ok, message)
"""
try:
info = model_info(repo_id, token=token)
pipeline_tag = getattr(info, "pipeline_tag", "unknown")
return True, f"✅ Found: `{repo_id}` (pipeline: {pipeline_tag})"
except RepositoryNotFoundError:
return False, f"❌ Repository not found: `{repo_id}`"
except HFValidationError as e:
return False, f"❌ Validation error: {e}"
except Exception as e:
return False, f"❌ Unexpected error: {e}"
def validate_url(url: str, timeout: int = 5) -> tuple[bool, str]:
"""HEAD-request an external model URL to check reachability."""
try:
r = requests.head(url, timeout=timeout, allow_redirects=True)
if r.status_code < 400:
return True, f"✅ URL reachable (HTTP {r.status_code})"
return False, f"❌ URL returned HTTP {r.status_code}"
except requests.ConnectionError:
return False, "❌ Cannot reach URL — check your connection"
except requests.Timeout:
return False, "❌ URL timed out"
except Exception as e:
return False, f"❌ {e}"
def validate_custom_model(model_id: str) -> tuple[bool, str]:
"""
Top-level validator: routes to HF or URL check depending on input shape.
Safe to call from Gradio event handlers.
"""
model_id = model_id.strip()
if not model_id:
return False, "⚠️ No model ID entered"
source = classify_model_id(model_id)
if source == "hf_repo":
return validate_hf_repo(model_id)
if source == "url":
return validate_url(model_id)
return False, f"❌ `{model_id}` does not look like a HF repo ID or a URL"
def search_hf_models(query: str, pipeline_tag: str | None = None, limit: int = 10) -> list[dict]:
"""
Search HuggingFace Hub for models matching a query.
Returns a list of dicts with keys: model_id, pipeline_tag, downloads, likes.
"""
try:
api = HfApi()
results = api.list_models(
search=query,
filter=pipeline_tag,
sort="downloads",
direction=-1,
limit=limit,
)
out = []
for m in results:
out.append({
"model_id": m.id,
"pipeline_tag": getattr(m, "pipeline_tag", "—"),
"downloads": getattr(m, "downloads", 0),
"likes": getattr(m, "likes", 0),
})
return out
except Exception as e:
logger.warning("HF search failed: %s", e)
return []