Spaces:
Sleeping
Sleeping
| """Model loading, unloading, status, and Hub search.""" | |
| from __future__ import annotations | |
| from collections.abc import Iterator | |
| import gradio as gr | |
| import torch | |
| from huggingface_hub import HfApi | |
| from api.auth import resolve_token | |
| from api.serialize import error_payload | |
| from miru_tracer.core.logging_config import get_logger | |
| from miru_tracer.core.model_manager import ModelManager | |
| from miru_tracer.core.session_manager import get_session_manager | |
| logger = get_logger(__name__) | |
| model_manager = ModelManager() | |
| QUICK_MODEL_CHOICES = ( | |
| "LiquidAI/LFM2.5-230M", | |
| "Qwen/Qwen3-0.6B", | |
| "Qwen/Qwen3-4B", | |
| ) | |
| def memory_usage() -> str: | |
| if torch.cuda.is_available(): | |
| allocated = torch.cuda.memory_allocated(0) / 1e9 | |
| total = torch.cuda.get_device_properties(0).total_memory / 1e9 | |
| percentage = (allocated / total) * 100 if total > 0 else 0 | |
| return f"{allocated:.2f} GB / {total:.2f} GB ({percentage:.1f}%)" | |
| return "CPU mode (no GPU)" | |
| def model_status() -> dict: | |
| """Current model state; also restores the Model view after a page load.""" | |
| info = None | |
| if model_manager.is_loaded(): | |
| model = model_manager.get_model() | |
| tokenizer = model_manager.get_tokenizer() | |
| info = { | |
| "model_name": model_manager.get_model_name(), | |
| "device": model_manager.get_device(), | |
| "vocab_size": len(tokenizer), | |
| "num_parameters_b": model.num_parameters() / 1e9, | |
| } | |
| return { | |
| "ok": True, | |
| "loaded": model_manager.is_loaded(), | |
| "model_name": model_manager.get_model_name(), | |
| "memory": memory_usage(), | |
| "cuda": torch.cuda.is_available(), | |
| "info": info, | |
| "quick_models": list(QUICK_MODEL_CHOICES), | |
| } | |
| def load_model( | |
| model_name: str, | |
| quantization: str, | |
| trust_remote_code: bool, | |
| minimize_ram: bool, | |
| oauth_token: gr.OAuthToken | None = None, | |
| ) -> Iterator[dict]: | |
| """Load a model, streaming status so the user sees progress. | |
| Uses the signed-in user's OAuth token for gated/private models, falling | |
| back to the Space's HF_TOKEN secret / local login. | |
| """ | |
| model_name = (model_name or "").strip() | |
| if not model_name: | |
| yield error_payload("Please enter a model name") | |
| return | |
| logger.info( | |
| f"Model load requested: {model_name} (quantization={quantization}, " | |
| f"trust_remote_code={trust_remote_code}, minimize_ram={minimize_ram})" | |
| ) | |
| if trust_remote_code: | |
| logger.warning("trust_remote_code=True enabled (security risk)") | |
| progress_lines = [f"Loading model: {model_name}"] | |
| if quantization != "none": | |
| progress_lines.append(f"Quantization: {quantization}") | |
| if trust_remote_code: | |
| progress_lines.append("Warning: trust_remote_code=True") | |
| if minimize_ram: | |
| progress_lines.append("RAM optimization: enabled (slower loading)") | |
| progress_lines.append("") | |
| progress_lines.append( | |
| "Downloading/loading weights — this can take a while for large " | |
| "models. Please wait..." | |
| ) | |
| yield {"ok": True, "done": False, "status": "\n".join(progress_lines)} | |
| try: | |
| _model, _tokenizer, _device, info = model_manager.load_model( | |
| model_name=model_name, | |
| quantization=quantization, | |
| trust_remote_code=bool(trust_remote_code), | |
| minimize_ram_usage=bool(minimize_ram), | |
| token=resolve_token(oauth_token), | |
| ) | |
| success_lines = [ | |
| "Model loaded successfully.", | |
| f"Device: {info['device_name']}", | |
| f"Vocabulary size: {info['vocab_size']:,}", | |
| f"Parameters: {info['num_parameters_b']:.2f}B", | |
| ] | |
| if info["device"] == "cuda": | |
| success_lines.append(f"VRAM: {info['vram_gb']:.2f} GB") | |
| if info.get("quantization_note"): | |
| success_lines.append(f"\n⚠️ {info['quantization_note']}") | |
| if info.get("is_vlm"): | |
| success_lines.append(f"\n⚠️ {info['vlm_warning']}") | |
| yield { | |
| "ok": True, | |
| "done": True, | |
| "status": "\n".join(success_lines), | |
| **model_status(), | |
| } | |
| except RuntimeError as e: | |
| # Concurrent load/unload in progress | |
| logger.warning(f"Load blocked: {e}") | |
| yield error_payload(str(e)) | |
| except Exception as e: | |
| logger.error(f"Model load failed: {model_name} - {e}", exc_info=True) | |
| hint = "" | |
| if "gated" in str(e).lower() or "401" in str(e) or "403" in str(e): | |
| hint = ( | |
| "\n\nThis looks like a gated or private model — sign in with " | |
| "Hugging Face (top right) with an account that has access." | |
| ) | |
| yield error_payload(f"Error loading model:\n\n{e}{hint}") | |
| def unload_model() -> dict: | |
| try: | |
| cleared_count = get_session_manager().clear_all_sessions() | |
| result = model_manager.unload_model() | |
| status = result["message"] | |
| if cleared_count: | |
| status += ( | |
| f"\n\n{cleared_count} active Interactive session(s) were " | |
| "cleared.\nAny in-progress generation work has been lost." | |
| ) | |
| if torch.cuda.is_available(): | |
| status += "\n\nGPU memory has been freed." | |
| status += ( | |
| "\n\nNote: any Generate-view run in progress is invalidated — " | |
| "start a new generation there." | |
| ) | |
| logger.info("Model unload completed") | |
| return {"ok": True, "status": status, **model_status()} | |
| except RuntimeError as e: | |
| logger.warning(f"Unload blocked: {e}") | |
| return error_payload(str(e)) | |
| def search_models(query: str, oauth_token: gr.OAuthToken | None = None) -> dict: | |
| """Search the Hub for text-generation models (powers the loader autocomplete).""" | |
| query = (query or "").strip() | |
| if len(query) < 2: | |
| return {"ok": True, "results": []} | |
| try: | |
| api = HfApi(token=resolve_token(oauth_token)) | |
| models = api.list_models( | |
| search=query, | |
| pipeline_tag="text-generation", | |
| sort="downloads", | |
| direction=-1, | |
| limit=12, | |
| ) | |
| results = [ | |
| { | |
| "id": m.id, | |
| "downloads": getattr(m, "downloads", None), | |
| "likes": getattr(m, "likes", None), | |
| "gated": bool(getattr(m, "gated", False)), | |
| } | |
| for m in models | |
| ] | |
| return {"ok": True, "results": results} | |
| except Exception as e: | |
| logger.warning(f"Model search failed: {e}") | |
| return {"ok": True, "results": [], "note": str(e)} | |