# SPDX-License-Identifier: AGPL-3.0-only - See /studio/LICENSE.AGPL-3.0 # Copyright © 2025 Unsloth AI """ Pydantic schemas for Model Management API """ from pydantic import BaseModel, Field from typing import Optional, List, Dict, Any, Literal ModelType = Literal["text", "vision", "audio", "embeddings"] class CheckpointInfo(BaseModel): """Information about a discovered checkpoint directory.""" display_name: str = Field(..., description="User-friendly checkpoint name (folder name)") path: str = Field(..., description="Full path to the checkpoint directory") loss: Optional[float] = Field(None, description="Training loss at this checkpoint") class ModelCheckpoints(BaseModel): """A training run and its associated checkpoints.""" name: str = Field(..., description="Training run folder name") checkpoints: List[CheckpointInfo] = Field( default_factory=list, description="List of checkpoints for this training run (final + intermediate)", ) base_model: Optional[str] = Field( None, description="Base model name from adapter_config.json or config.json", ) peft_type: Optional[str] = Field( None, description="PEFT type (e.g. LORA) if adapter training, None for full fine-tune", ) lora_rank: Optional[int] = Field( None, description="LoRA rank (r) if applicable", ) class CheckpointListResponse(BaseModel): """Response for listing available checkpoints in an outputs directory.""" outputs_dir: str = Field(..., description="Directory that was scanned") models: List[ModelCheckpoints] = Field( default_factory=list, description="List of training runs with their checkpoints", ) class ModelDetails(BaseModel): """Detailed model configuration and metadata - can be used for both list and detail views""" id: str = Field(..., description="Model identifier") model_name: Optional[str] = Field(None, description="Model identifier (alias for id, for backward compatibility)") name: Optional[str] = Field(None, description="Display name for the model") config: Optional[Dict[str, Any]] = Field(None, description="Model configuration dictionary") is_vision: bool = Field(False, description="Whether model is a vision model") is_embedding: bool = Field(False, description="Whether model is an embedding/sentence-transformer 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 format)") is_audio: bool = Field(False, description="Whether model is a TTS audio model") audio_type: Optional[str] = Field(None, description="Audio codec type: snac, csm, bicodec, dac") has_audio_input: bool = Field(False, description="Whether model accepts audio input (ASR)") model_type: Optional[ModelType] = Field(None, description="Collapsed model modality: text, vision, audio, or embeddings") base_model: Optional[str] = Field(None, description="Base model if this is a LoRA adapter") class LoRAInfo(BaseModel): """LoRA adapter or exported model information""" display_name: str = Field(..., description="Display name for the LoRA") adapter_path: str = Field(..., description="Path to the LoRA adapter or exported model") base_model: Optional[str] = Field(None, description="Base model identifier") source: Optional[str] = Field(None, description="'training' or 'exported'") export_type: Optional[str] = Field(None, description="'lora', 'merged', or 'gguf' (for exports)") class LoRAScanResponse(BaseModel): """Response schema for scanning trained LoRA adapters""" loras: List[LoRAInfo] = Field(default_factory=list, description="List of found LoRA adapters") outputs_dir: str = Field(..., description="Directory that was scanned") class ModelListResponse(BaseModel): """Response schema for listing models""" models: List[ModelDetails] = Field(default_factory=list, description="List of models") default_models: List[str] = Field(default_factory=list, description="List of default model IDs") class GgufVariantDetail(BaseModel): """A single GGUF quantization variant in a HuggingFace repo.""" filename: str = Field(..., description="GGUF filename (e.g., 'gemma-3-4b-it-Q4_K_M.gguf')") quant: str = Field(..., description="Quantization label (e.g., 'Q4_K_M')") size_bytes: int = Field(0, description="File size in bytes") class GgufVariantsResponse(BaseModel): """Response for listing GGUF quantization variants in a HuggingFace repo.""" repo_id: str = Field(..., description="HuggingFace repo ID") variants: List[GgufVariantDetail] = Field(default_factory=list, description="Available GGUF variants") has_vision: bool = Field(False, description="Whether the model has vision support (mmproj files)") default_variant: Optional[str] = Field(None, description="Recommended default quantization variant") class LocalModelInfo(BaseModel): """Discovered local model candidate.""" id: str = Field(..., description="Identifier to use for loading/training") display_name: str = Field(..., description="Display label") path: str = Field(..., description="Local path where model data was discovered") source: Literal["models_dir", "hf_cache"] = Field( ..., description="Discovery source", ) model_id: Optional[str] = Field( None, description="HF repo id for cached models, e.g. org/model", ) updated_at: Optional[float] = Field( None, description="Unix timestamp of latest observed update", ) class LocalModelListResponse(BaseModel): """Response schema for listing local/cached models.""" models_dir: str = Field(..., description="Directory scanned for custom local models") hf_cache_dir: Optional[str] = Field( None, description="HF cache root that was scanned", ) models: List[LocalModelInfo] = Field( default_factory=list, description="Discovered local/cached models", )