Merge pull request #14 from unslothai/feature/pydantic-models-update

Feature/pydantic models update
This commit is contained in:
Roland Tannous 2026-02-03 14:47:12 +04:00 committed by GitHub
commit d8183237f3
3 changed files with 43 additions and 65 deletions

2
.gitignore vendored
View file

@ -20,7 +20,6 @@ unsloth_compiled_cache/
outputs/
*.gguf
*.safetensors
/models/
# IDE / Editors
.vscode/
@ -35,3 +34,4 @@ Thumbs.db
# Other
resources/
tmp/

View file

@ -5,36 +5,8 @@ from pydantic import BaseModel, Field
from typing import Optional, List, Dict, Any
class ModelSearchRequest(BaseModel):
"""Request schema for searching HuggingFace models"""
query: str = Field(..., description="Search query")
hf_token: Optional[str] = Field(None, description="HuggingFace token for authenticated searches")
class ModelInfo(BaseModel):
"""Model information"""
id: str = Field(..., description="Model identifier")
name: Optional[str] = Field(None, description="Display name")
description: Optional[str] = Field(None, description="Model description")
size: Optional[str] = Field(None, description="Model size")
is_vision: bool = Field(False, description="Whether model is a vision model")
is_lora: bool = Field(False, description="Whether model is a LoRA adapter")
class ModelSearchResponse(BaseModel):
"""Response schema for model search"""
models: List[ModelInfo] = Field(default_factory=list, description="List of matching models")
total: int = Field(0, description="Total number of results")
class ModelListResponse(BaseModel):
"""Response schema for listing available models"""
models: List[ModelInfo] = Field(default_factory=list, description="List of available models")
default_models: List[str] = Field(default_factory=list, description="List of default model IDs")
class ModelConfigResponse(BaseModel):
"""Response schema for model configuration"""
class ModelDetails(BaseModel):
"""Detailed model configuration and metadata"""
model_name: str = Field(..., description="Model identifier")
config: Dict[str, Any] = Field(..., description="Model configuration dictionary")
is_vision: bool = Field(False, description="Whether model is a vision model")

View file

@ -2,7 +2,7 @@
Pydantic schemas for Training API
"""
from pydantic import BaseModel, Field
from typing import Optional, List
from typing import Optional, List, Literal
class TrainingStartRequest(BaseModel):
@ -59,38 +59,44 @@ class TrainingStartRequest(BaseModel):
tensorboard_dir: Optional[str] = Field(None, description="TensorBoard directory")
class TrainingStartResponse(BaseModel):
"""Response schema for training start"""
status: str = Field(..., description="Status: 'started' or 'error'")
job_id: Optional[str] = Field(None, description="Training job ID")
message: str = Field(..., description="Status message")
error: Optional[str] = Field(None, description="Error message if status is 'error'")
class TrainingJobResponse(BaseModel):
"""Immediate response when training is initiated"""
job_id: str = Field(..., description="Unique training job identifier")
status: Literal["queued", "error"] = Field(..., description="Initial job status")
message: str = Field(..., description="Human-readable status message")
error: Optional[str] = Field(None, description="Error details if status is 'error'")
class TrainingStatusResponse(BaseModel):
"""Response schema for training status"""
status: str = Field(..., description="Status: 'idle', 'preparing', 'training', 'stopping', 'error'")
is_active: bool = Field(..., description="Whether training is currently active (actual training running)")
message: str = Field(..., description="Status message")
current_step: Optional[int] = Field(None, description="Current training step")
total_steps: Optional[int] = Field(None, description="Total training steps")
class TrainingStatus(BaseModel):
"""Current training job status - works for streaming or polling"""
job_id: str = Field(..., description="Training job identifier")
phase: Literal[
"idle",
"loading_model",
"loading_dataset",
"configuring",
"training",
"completed",
"error",
"stopped"
] = Field(..., description="Current phase of training pipeline")
is_training_running: bool = Field(..., description="True if training loop is actively running")
message: str = Field(..., description="Human-readable status message")
error: Optional[str] = Field(None, description="Error details if phase is 'error'")
details: Optional[dict] = Field(None, description="Phase-specific info, e.g. {'model_size': '8B'}")
class TrainingMetricsResponse(BaseModel):
"""Response schema for training metrics"""
loss_history: List[float] = Field(default_factory=list, description="Loss values")
lr_history: List[float] = Field(default_factory=list, description="Learning rate values")
step_history: List[int] = Field(default_factory=list, description="Step numbers")
current_loss: Optional[float] = Field(None, description="Current loss value")
current_lr: Optional[float] = Field(None, description="Current learning rate")
current_step: Optional[int] = Field(None, description="Current step")
class TrainingProgressResponse(BaseModel):
"""Response schema for training progress updates"""
step: int = Field(..., description="Current step")
loss: float = Field(..., description="Current loss")
class TrainingProgress(BaseModel):
"""Training progress metrics - for streaming or polling"""
job_id: str = Field(..., description="Training job identifier")
step: int = Field(..., description="Current training step")
total_steps: int = Field(..., description="Total training steps")
loss: float = Field(..., description="Current loss value")
learning_rate: float = Field(..., description="Current learning rate")
status_message: str = Field(..., description="Status message")
progress_percent: Optional[float] = Field(None, description="Progress percentage")
progress_percent: float = Field(..., description="Progress percentage (0.0 to 100.0)")
epoch: Optional[int] = Field(None, description="Current epoch")
elapsed_seconds: Optional[float] = Field(None, description="Time elapsed since training started")
eta_seconds: Optional[float] = Field(None, description="Estimated time remaining")
grad_norm: Optional[float] = Field(None, description="L2 norm of gradients, computed before gradient clipping")
num_tokens: Optional[int] = Field(None, description="Total number of tokens processed so far")