Merge pull request #14 from unslothai/feature/pydantic-models-update
Feature/pydantic models update
This commit is contained in:
commit
d8183237f3
3 changed files with 43 additions and 65 deletions
2
.gitignore
vendored
2
.gitignore
vendored
|
|
@ -20,7 +20,6 @@ unsloth_compiled_cache/
|
|||
outputs/
|
||||
*.gguf
|
||||
*.safetensors
|
||||
/models/
|
||||
|
||||
# IDE / Editors
|
||||
.vscode/
|
||||
|
|
@ -35,3 +34,4 @@ Thumbs.db
|
|||
|
||||
# Other
|
||||
resources/
|
||||
tmp/
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue