unsloth/studio/backend/models/training.py
Roland Tannous 202780c32c feat: Dataset Conversion Advisor — multi-pass LLM for non-conversational datasets
Non-conversational HF datasets (e.g. stanfordnlp/snli) were naively mapped
column→role, producing poor training results. The AI Assist button now runs
a 3-pass advisor using Qwen 7B that:
1. Fetches the HF dataset card/README to understand the dataset purpose
2. Classifies the dataset type and determines if conversion is needed
3. Generates a system prompt, user/assistant templates with {column}
   placeholders, and label mappings (e.g. 0→entailment)
4. Validates the conversion quality (score ≥7/10 required)

Architecture: advisor metadata flows as __-prefixed keys in
custom_format_mapping (e.g. __system_prompt, __user_template,
__assistant_template, __label_mapping). The existing _apply_user_mapping()
detects these keys and routes to template-based conversation construction.
No __ keys = existing simple mode (backwards compatible).

Backend: upgraded llm_assist.py (7B default, multi-pass advisor,
HF card fetching), extended API models, added _apply_template_mapping()
to dataset_utils.py.

Frontend: extended store with advisor state fields, wired AI Assist
to store templates/system prompt, inject __ metadata in training request,
show advisor notification banner in mapping card.
2026-03-10 15:39:56 +00:00

139 lines
7.8 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only - See /studio/LICENSE.AGPL-3.0
# Copyright © 2025 Unsloth AI
"""
Pydantic schemas for Training API
"""
from pydantic import BaseModel, Field, model_validator
from typing import Any, Optional, List, Dict, Literal
class TrainingStartRequest(BaseModel):
"""Request schema for starting training"""
# Model parameters
model_name: str = Field(..., description="Model identifier (e.g., 'unsloth/llama-3-8b-bnb-4bit')")
training_type: str = Field(..., description="Training type: 'LoRA/QLoRA' or 'Full Finetuning'")
hf_token: Optional[str] = Field(None, description="HuggingFace token")
load_in_4bit: bool = Field(True, description="Load model in 4-bit quantization")
max_seq_length: int = Field(2048, description="Maximum sequence length")
trust_remote_code: bool = Field(
False,
description="Allow loading models with custom code (e.g. NVIDIA Nemotron). Only enable for repos you trust.",
)
# Dataset parameters
hf_dataset: Optional[str] = Field(None, description="HuggingFace dataset identifier")
local_datasets: List[str] = Field(default_factory=list, description="List of local dataset paths")
format_type: str = Field(..., description="Dataset format type")
subset: Optional[str] = None
train_split: Optional[str] = Field("train", description="Training split name")
eval_split: Optional[str] = Field(None, description="Eval split name. None = auto-detect")
eval_steps: float = Field(0.00, description="Fraction of total steps between evals (0-1)")
dataset_slice_start: Optional[int] = Field(None, description="Inclusive start row index for dataset slicing")
dataset_slice_end: Optional[int] = Field(None, description="Inclusive end row index for dataset slicing")
@model_validator(mode="before")
@classmethod
def _compat_split(cls, values: Any) -> Any:
"""Accept legacy 'split' field as alias for 'train_split'."""
if isinstance(values, dict) and "split" in values:
values.setdefault("train_split", values.pop("split"))
return values
custom_format_mapping: Optional[Dict[str, Any]] = Field(
None,
description=(
"User-provided column-to-role mapping, e.g. {'image': 'image', 'caption': 'text'} "
"for VLM or {'instruction': 'user', 'output': 'assistant'} for LLM. "
"Enhanced format includes __system_prompt, __user_template, "
"__assistant_template, __label_mapping metadata keys."
),
)
# Training parameters
num_epochs: int = Field(1, description="Number of training epochs")
learning_rate: str = Field("2e-4", description="Learning rate")
batch_size: int = Field(1, description="Batch size")
gradient_accumulation_steps: int = Field(1, description="Gradient accumulation steps")
warmup_steps: Optional[int] = Field(None, description="Warmup steps")
warmup_ratio: Optional[float] = Field(None, description="Warmup ratio")
max_steps: Optional[int] = Field(None, description="Maximum training steps")
save_steps: int = Field(100, description="Steps between checkpoints")
weight_decay: float = Field(0.01, description="Weight decay")
random_seed: int = Field(42, description="Random seed")
packing: bool = Field(False, description="Enable sequence packing")
optim: str = Field("adamw_8bit", description="Optimizer")
lr_scheduler_type: str = Field("linear", description="Learning rate scheduler type")
# LoRA parameters
use_lora: bool = Field(True, description="Use LoRA (derived from training_type)")
lora_r: int = Field(16, description="LoRA rank")
lora_alpha: int = Field(16, description="LoRA alpha")
lora_dropout: float = Field(0.0, description="LoRA dropout")
target_modules: List[str] = Field(default_factory=list, description="Target modules for LoRA")
gradient_checkpointing: str = Field("", description="Gradient checkpointing setting")
use_rslora: bool = Field(False, description="Use RSLoRA")
use_loftq: bool = Field(False, description="Use LoftQ")
train_on_completions: bool = Field(False, description="Train on completions only")
# Vision-specific LoRA parameters
finetune_vision_layers: bool = Field(False, description="Finetune vision layers")
finetune_language_layers: bool = Field(False, description="Finetune language layers")
finetune_attention_modules: bool = Field(False, description="Finetune attention modules")
finetune_mlp_modules: bool = Field(False, description="Finetune MLP modules")
is_dataset_image: bool = Field(False, description="Whether the dataset contains image data")
is_dataset_audio: bool = Field(False, description="Whether the dataset contains audio data")
# Logging parameters
enable_wandb: bool = Field(False, description="Enable Weights & Biases logging")
wandb_token: Optional[str] = Field(None, description="W&B token")
wandb_project: Optional[str] = Field(None, description="W&B project name")
enable_tensorboard: bool = Field(False, description="Enable TensorBoard logging")
tensorboard_dir: Optional[str] = Field(None, description="TensorBoard directory")
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 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")
eval_enabled: bool = Field(False, description="True if evaluation dataset is configured for this training run")
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'}")
metric_history: Optional[dict] = Field(
None,
description="Full metric history arrays for chart recovery after SSE reconnection. "
"Keys: 'steps', 'loss', 'lr', 'grad_norm', 'grad_norm_steps' — each a list of numeric values.",
)
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")
progress_percent: float = Field(..., description="Progress percentage (0.0 to 100.0)")
epoch: Optional[float] = 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")
eval_loss: Optional[float] = Field(None, description="Eval loss from the most recent evaluation step")