Refactored the training and model routes and added the jwt authentication

This commit is contained in:
sshah229 2026-02-04 05:49:29 -07:00
commit 009e93f079
6 changed files with 393 additions and 161 deletions

View file

@ -0,0 +1,69 @@
import secrets
from datetime import UTC, datetime, timedelta
from typing import Optional
from fastapi import Depends, HTTPException, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from jose import JWTError, jwt
# Ephemeral in-memory secret:
# - Generated fresh on each backend process start
# - Never written to disk
# - Not configurable by the user
SECRET_KEY = secrets.token_urlsafe(64)
ALGORITHM = "HS256"
ACCESS_TOKEN_EXPIRE_MINUTES = 60
security = HTTPBearer() # Reads Authorization: Bearer <token>
def create_access_token(
subject: str,
expires_delta: Optional[timedelta] = None,
) -> str:
"""
Create a signed JWT for the given subject (e.g. "local-user").
Tokens are valid only for the lifetime of this process, because the
SECRET_KEY is regenerated each time the backend restarts.
"""
to_encode = {"sub": subject}
expire = datetime.now(UTC) + (
expires_delta or timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
)
to_encode.update({"exp": expire})
return jwt.encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)
async def get_current_subject(
credentials: HTTPAuthorizationCredentials = Depends(security),
) -> str:
"""
FastAPI dependency to validate the JWT and return the subject.
Use this as a dependency on routes that should be protected, e.g.:
@router.get("/secure")
async def secure_endpoint(current_subject: str = Depends(get_current_subject)):
...
"""
token = credentials.credentials
try:
payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
subject: Optional[str] = payload.get("sub")
if subject is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Invalid token payload",
)
return subject
except JWTError:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Invalid or expired token",
)
token = create_access_token("local-user")
print(token)

View file

@ -3,10 +3,10 @@ Unified core module for Unsloth backend
"""
# Inference
from .inference import InferenceBackend, get_inference_backend
from .inference.inference import InferenceBackend, get_inference_backend
# Training
from .training import UnslothTrainer, get_trainer, TrainingBackend, get_training_backend, create_training_handlers, TrainingProgress
from .training.training import UnslothTrainer, get_trainer, TrainingBackend, get_training_backend, create_training_handlers, TrainingProgress
# Configuration (from utils)
from utils.models import is_vision_model, ModelConfig, scan_trained_loras, load_model_defaults, get_base_model_from_lora

View file

@ -3,35 +3,25 @@ Pydantic models for API request/response schemas
"""
from .training import (
TrainingStartRequest,
TrainingStartResponse,
TrainingStatusResponse,
TrainingMetricsResponse,
TrainingProgressResponse,
TrainingJobResponse,
TrainingStatus,
TrainingProgress,
)
from .models import (
ModelSearchRequest,
ModelSearchResponse,
ModelListResponse,
ModelConfigResponse,
LoRAScanResponse,
ModelDetails,
LoRAInfo,
ModelInfo,
LoRAScanResponse,
)
__all__ = [
# Training schemas
"TrainingStartRequest",
"TrainingStartResponse",
"TrainingStatusResponse",
"TrainingMetricsResponse",
"TrainingProgressResponse",
"TrainingJobResponse",
"TrainingStatus",
"TrainingProgress",
# Model management schemas
"ModelSearchRequest",
"ModelSearchResponse",
"ModelListResponse",
"ModelConfigResponse",
"LoRAScanResponse",
"ModelDetails",
"LoRAInfo",
"ModelInfo",
"LoRAScanResponse",
]

View file

@ -3,15 +3,19 @@ Model Management API routes
"""
import sys
from pathlib import Path
from fastapi import APIRouter, HTTPException, Query
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from typing import List, Optional
import logging
from pydantic import BaseModel
# Add backend directory to path
backend_path = Path(__file__).parent.parent.parent
if str(backend_path) not in sys.path:
sys.path.insert(0, str(backend_path))
from auth.jwt import get_current_subject
# Import backend functions
try:
from utils.utils import search_hf_models
@ -38,16 +42,42 @@ except ImportError:
)
from core.inference import get_inference_backend
from models.models import (
ModelSearchRequest,
ModelSearchResponse,
ModelInfo,
ModelListResponse,
ModelConfigResponse,
from models import (
ModelDetails,
LoRAScanResponse,
LoRAInfo,
)
class ModelInfo(BaseModel):
"""Basic model info used in search/list responses"""
id: str
name: Optional[str] = None
is_vision: Optional[bool] = False
is_lora: Optional[bool] = False
class ModelSearchRequest(BaseModel):
"""Request body for model search"""
query: str
hf_token: Optional[str] = None
class ModelSearchResponse(BaseModel):
"""Response schema for model search"""
models: List[ModelInfo]
total: int
class ModelListResponse(BaseModel):
"""Response schema for listing models"""
models: List[ModelInfo]
default_models: List[str]
router = APIRouter()
logger = logging.getLogger(__name__)
@ -62,7 +92,10 @@ if not logger.handlers:
@router.post("/search")
async def search_models(request: ModelSearchRequest):
async def search_models(
request: ModelSearchRequest,
current_subject: str = Depends(get_current_subject),
):
"""
Search for models on HuggingFace Hub.
@ -117,7 +150,9 @@ async def search_models(request: ModelSearchRequest):
@router.get("/list")
async def list_models():
async def list_models(
current_subject: str = Depends(get_current_subject),
):
"""
List available models (default models and loaded models).
@ -174,7 +209,10 @@ async def list_models():
@router.get("/config/{model_name:path}")
async def get_model_config(model_name: str):
async def get_model_config(
model_name: str,
current_subject: str = Depends(get_current_subject),
):
"""
Get configuration for a specific model.
@ -200,12 +238,12 @@ async def get_model_config(model_name: str):
# If ModelConfig creation fails, use defaults
pass
return ModelConfigResponse(
return ModelDetails(
model_name=model_name,
config=config_dict,
is_vision=is_vision,
is_lora=is_lora,
base_model=base_model
base_model=base_model,
)
except Exception as e:
@ -218,7 +256,8 @@ async def get_model_config(model_name: str):
@router.get("/loras")
async def scan_loras(
outputs_dir: str = Query(default="./outputs", description="Directory to scan for LoRA adapters")
outputs_dir: str = Query(default="./outputs", description="Directory to scan for LoRA adapters"),
current_subject: str = Depends(get_current_subject),
):
"""
Scan for trained LoRA adapters in the outputs directory.
@ -256,7 +295,10 @@ async def scan_loras(
@router.get("/loras/{lora_path:path}/base-model")
async def get_lora_base_model(lora_path: str):
async def get_lora_base_model(
lora_path: str,
current_subject: str = Depends(get_current_subject),
):
"""
Get the base model for a LoRA adapter.
@ -287,7 +329,10 @@ async def get_lora_base_model(lora_path: str):
@router.get("/check-vision/{model_name:path}")
async def check_vision_model(model_name: str):
async def check_vision_model(
model_name: str,
current_subject: str = Depends(get_current_subject),
):
"""
Check if a model is a vision model.

View file

@ -3,9 +3,9 @@ Training API routes
"""
import sys
from pathlib import Path
from fastapi import APIRouter, HTTPException
from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import StreamingResponse
from typing import Dict
from typing import Dict, Optional
import logging
import asyncio
from datetime import datetime
@ -27,12 +27,14 @@ except ImportError:
sys.path.insert(0, str(parent_backend))
from core.training import get_training_backend
from models.training import (
# Auth
from auth.jwt import get_current_subject
from models import (
TrainingStartRequest,
TrainingStartResponse,
TrainingStatusResponse,
TrainingMetricsResponse,
TrainingProgressResponse,
TrainingJobResponse,
TrainingStatus,
TrainingProgress,
)
router = APIRouter()
@ -49,35 +51,47 @@ if not logger.handlers:
@router.post("/start")
async def start_training(request: TrainingStartRequest):
async def start_training(
request: TrainingStartRequest,
current_subject: str = Depends(get_current_subject),
):
"""
Start a training job.
This endpoint initiates training in the background and returns immediately.
Use the /status endpoint to check training progress.
"""
try:
logger.info(f"Starting training job with model: {request.model_name}")
backend = get_training_backend()
# Generate job ID and attach to backend for later status/progress calls
job_id = f"job_{datetime.now().strftime('%Y%m%d_%H%M%S')}"
backend.current_job_id = job_id
# Check if training is already active
if backend.is_training_active():
return TrainingStartResponse(
existing_job_id: Optional[str] = getattr(backend, "current_job_id", "")
return TrainingJobResponse(
job_id=existing_job_id or job_id,
status="error",
message="Training is already in progress. Stop current training before starting a new one.",
error="Training already active"
message=(
"Training is already in progress. "
"Stop current training before starting a new one."
),
error="Training already active",
)
# Validate dataset paths if provided
if request.local_datasets:
validated_datasets = []
# Get the backend directory (where this file is located)
backend_dir = Path(__file__).parent.parent
assets_datasets_dir = backend_dir / "assets" / "datasets"
for dataset_path in request.local_datasets:
dataset_file = Path(dataset_path)
# If not absolute, try multiple locations
if not dataset_file.is_absolute():
# First try: relative to current working directory
@ -89,14 +103,16 @@ async def start_training(request: TrainingStartRequest):
# Third try: just the filename in assets/datasets
candidate = assets_datasets_dir / dataset_file.name
dataset_file = candidate
if not dataset_file.exists():
logger.warning(f"Dataset file not found: {dataset_path} (resolved: {dataset_file})")
logger.warning(
f"Dataset file not found: {dataset_path} (resolved: {dataset_file})"
)
else:
logger.info(f"Found dataset file: {dataset_file}")
validated_datasets.append(str(dataset_file))
request.local_datasets = validated_datasets
# Convert request to kwargs for backend
training_kwargs = {
"model_name": request.model_name,
@ -125,7 +141,9 @@ async def start_training(request: TrainingStartRequest):
"lora_alpha": request.lora_alpha,
"lora_dropout": request.lora_dropout,
"target_modules": request.target_modules if request.target_modules else None,
"gradient_checkpointing": request.gradient_checkpointing.strip() if request.gradient_checkpointing and request.gradient_checkpointing.strip() else "unsloth",
"gradient_checkpointing": request.gradient_checkpointing.strip()
if request.gradient_checkpointing and request.gradient_checkpointing.strip()
else "unsloth",
"use_rslora": request.use_rslora,
"use_loftq": request.use_loftq,
"train_on_completions": request.train_on_completions,
@ -139,84 +157,95 @@ async def start_training(request: TrainingStartRequest):
"enable_tensorboard": request.enable_tensorboard,
"tensorboard_dir": request.tensorboard_dir or "",
}
# Generate job ID
job_id = f"job_{datetime.now().strftime('%Y%m%d_%H%M%S')}"
# Set initial "preparing" state
try:
backend.trainer._update_progress(
status_message="Initializing training...",
is_training=False
is_training=False,
)
except:
except Exception:
pass
def run_training():
try:
logger.info(f"Starting training job {job_id} with model {request.model_name}")
logger.info(
f"Starting training job {job_id} with model {request.model_name}"
)
# Update status to show we're loading model
try:
backend.trainer._update_progress(status_message="Loading model...")
except Exception as e:
logger.error(f"Error updating progress: {e}")
# Consume the generator - this actually runs the training
update_count = 0
for update_tuple in backend.start_training(**training_kwargs):
for _update_tuple in backend.start_training(**training_kwargs):
update_count += 1
if update_count % 10 == 0:
logger.info(f"Training progress update #{update_count}")
logger.info(f"Training job {job_id} completed successfully")
except Exception as e:
logger.error(f"Training error in job {job_id}: {e}", exc_info=True)
try:
backend.trainer._update_progress(
error=str(e),
is_training=False
is_training=False,
)
except Exception as update_error:
logger.error(f"Failed to update progress: {update_error}")
# Start training in a daemon thread
training_thread = threading.Thread(target=run_training, daemon=True, name=f"Training-{job_id}")
training_thread = threading.Thread(
target=run_training,
daemon=True,
name=f"Training-{job_id}",
)
training_thread.start()
# Store thread reference for status checking
backend._training_thread = training_thread
# Give it a moment to start
import time
time.sleep(0.5)
# Verify training thread is alive
if not training_thread.is_alive():
logger.warning(f"Training thread died immediately for job {job_id}")
return TrainingStartResponse(
return TrainingJobResponse(
job_id=job_id,
status="error",
message="Training thread failed to start. Check server logs for details.",
error="Thread not alive"
message=(
"Training thread failed to start. "
"Check server logs for details."
),
error="Thread not alive",
)
return TrainingStartResponse(
status="started",
return TrainingJobResponse(
job_id=job_id,
message="Training job started successfully"
status="queued",
message="Training job queued and starting in background",
error=None,
)
except Exception as e:
logger.error(f"Error starting training: {e}", exc_info=True)
raise HTTPException(
status_code=500,
detail=f"Failed to start training: {str(e)}"
detail=f"Failed to start training: {str(e)}",
)
@router.post("/stop")
async def stop_training():
async def stop_training(
current_subject: str = Depends(get_current_subject),
):
"""
Stop the currently running training job.
"""
@ -246,55 +275,75 @@ async def stop_training():
@router.get("/status")
async def get_training_status():
async def get_training_status(
current_subject: str = Depends(get_current_subject),
):
"""
Get the current training status.
"""
try:
backend = get_training_backend()
job_id: str = getattr(backend, "current_job_id", "")
# Check if training is active
is_active = backend.is_training_active()
# Check if there's a training thread running (preparation phase)
has_thread = hasattr(backend, '_training_thread') and backend._training_thread and backend._training_thread.is_alive()
# Get progress info
has_thread = (
hasattr(backend, "_training_thread")
and backend._training_thread
and backend._training_thread.is_alive()
)
# Get progress info from trainer
try:
progress = backend.trainer.get_training_progress()
status_message = progress.status_message or "Ready to train"
except:
except Exception:
progress = None
status_message = "Unknown"
if is_active:
# Actual training is running
trainer = backend.trainer
current_step = getattr(trainer.training_progress, 'step', None) or (progress.step if progress else None)
total_steps = getattr(trainer.training_progress, 'total_steps', None) or (progress.total_steps if progress else None)
return TrainingStatusResponse(
status="training",
is_active=True,
message=status_message or "Training is in progress",
current_step=current_step,
total_steps=total_steps
)
elif has_thread or (progress and status_message and any(keyword in status_message.lower() for keyword in ["loading", "preparing", "initializing"])):
# Training thread is running but not yet in active training phase
return TrainingStatusResponse(
status="preparing",
is_active=False,
message=status_message or "Preparing training...",
current_step=None,
total_steps=None
)
status_message = (
getattr(progress, "status_message", None) if progress else None
) or "Ready to train"
error_message = getattr(progress, "error", None) if progress else None
# Derive high-level phase
if error_message:
phase = "error"
elif is_active:
msg_lower = status_message.lower()
if "loading" in msg_lower:
phase = "loading_model"
elif any(
k in msg_lower for k in ["preparing", "initializing", "configuring"]
):
phase = "configuring"
else:
phase = "training"
elif progress and getattr(progress, "is_completed", False):
phase = "completed"
elif has_thread:
phase = "loading_model"
else:
return TrainingStatusResponse(
status="idle",
is_active=False,
message="No training job is currently running"
)
phase = "idle"
details = None
if progress:
details = {
"epoch": getattr(progress, "epoch", 0),
"step": getattr(progress, "step", 0),
"total_steps": getattr(progress, "total_steps", 0),
"loss": getattr(progress, "loss", 0.0),
"learning_rate": getattr(progress, "learning_rate", 0.0),
}
return TrainingStatus(
job_id=job_id,
phase=phase,
is_training_running=is_active,
message=status_message,
error=error_message,
details=details,
)
except Exception as e:
logger.error(f"Error getting training status: {e}", exc_info=True)
@ -305,7 +354,9 @@ async def get_training_status():
@router.get("/metrics")
async def get_training_metrics():
async def get_training_metrics(
current_subject: str = Depends(get_current_subject),
):
"""
Get training metrics (loss, learning rate, steps).
"""
@ -316,20 +367,21 @@ async def get_training_metrics():
loss_history = backend.loss_history
lr_history = backend.lr_history
step_history = backend.step_history
# Get current values
current_loss = loss_history[-1] if loss_history else None
current_lr = lr_history[-1] if lr_history else None
current_step = step_history[-1] if step_history else None
return TrainingMetricsResponse(
loss_history=loss_history,
lr_history=lr_history,
step_history=step_history,
current_loss=current_loss,
current_lr=current_lr,
current_step=current_step
)
# Keep metrics as a simple JSON payload instead of a Pydantic model
return {
"loss_history": loss_history,
"lr_history": lr_history,
"step_history": step_history,
"current_loss": current_loss,
"current_lr": current_lr,
"current_step": current_step,
}
except Exception as e:
logger.error(f"Error getting training metrics: {e}", exc_info=True)
@ -340,7 +392,9 @@ async def get_training_metrics():
@router.get("/progress")
async def stream_training_progress():
async def stream_training_progress(
current_subject: str = Depends(get_current_subject),
):
"""
Stream training progress updates using Server-Sent Events (SSE).
@ -348,12 +402,53 @@ async def stream_training_progress():
"""
async def event_generator():
backend = get_training_backend()
job_id: str = getattr(backend, "current_job_id", "")
# Helper to build a TrainingProgress payload from raw values
def build_progress(
step: int,
loss: float,
learning_rate: float,
total_steps: int,
epoch: Optional[int] = None,
) -> TrainingProgress:
total = max(total_steps, 0)
if step < 0 or total == 0:
progress_percent = 0.0
else:
progress_percent = (
float(step) / float(total) * 100.0 if total > 0 else 0.0
)
return TrainingProgress(
job_id=job_id,
step=step,
total_steps=total,
loss=loss,
learning_rate=learning_rate,
progress_percent=progress_percent,
epoch=epoch,
elapsed_seconds=None,
eta_seconds=None,
grad_norm=None,
num_tokens=None,
)
# Send initial status
is_active = backend.is_training_active()
initial_message = 'Connecting...' if is_active else 'No training in progress'
yield f"data: {TrainingProgressResponse(step=0, loss=0.0, learning_rate=0.0, status_message=initial_message).model_dump_json()}\n\n"
tp = getattr(getattr(backend, "trainer", None), "training_progress", None)
initial_total_steps = getattr(tp, "total_steps", 0) if tp else 0
initial_epoch = getattr(tp, "epoch", None) if tp else None
initial_progress = build_progress(
step=0,
loss=0.0,
learning_rate=0.0,
total_steps=initial_total_steps,
epoch=initial_epoch,
)
yield f"data: {initial_progress.model_dump_json()}\n\n"
# If not active, check if there's any history
if not is_active:
if backend.step_history:
@ -361,9 +456,13 @@ async def stream_training_progress():
final_step = backend.step_history[-1]
final_loss = backend.loss_history[-1] if backend.loss_history else 0.0
final_lr = backend.lr_history[-1] if backend.lr_history else 0.0
yield f"data: {TrainingProgressResponse(step=final_step, loss=final_loss, learning_rate=final_lr, status_message='Training completed').model_dump_json()}\n\n"
final_total_steps = (
getattr(tp, "total_steps", final_step) if tp else final_step
)
final_epoch = getattr(tp, "epoch", None) if tp else None
yield f"data: {build_progress(final_step, final_loss, final_lr, final_total_steps, final_epoch).model_dump_json()}\n\n"
else:
yield f"data: {TrainingProgressResponse(step=-1, loss=0.0, learning_rate=0.0, status_message='No training in progress').model_dump_json()}\n\n"
yield f"data: {build_progress(-1, 0.0, 0.0, 0).model_dump_json()}\n\n"
return
# Poll for updates while training is active
@ -378,53 +477,81 @@ async def stream_training_progress():
current_step = backend.step_history[-1]
current_loss = backend.loss_history[-1] if backend.loss_history else 0.0
current_lr = backend.lr_history[-1] if backend.lr_history else 0.0
tp_inner = getattr(
getattr(backend, "trainer", None), "training_progress", None
)
current_total_steps = (
getattr(tp_inner, "total_steps", current_step)
if tp_inner
else current_step
)
current_epoch = getattr(tp_inner, "epoch", None) if tp_inner else None
# Only send if step changed
if current_step != last_step:
progress = TrainingProgressResponse(
step=current_step,
loss=current_loss,
learning_rate=current_lr,
status_message=f"Training step {current_step}"
progress_payload = build_progress(
current_step,
current_loss,
current_lr,
current_total_steps,
current_epoch,
)
yield f"data: {progress.model_dump_json()}\n\n"
yield f"data: {progress_payload.model_dump_json()}\n\n"
last_step = current_step
no_update_count = 0
else:
no_update_count += 1
# Send heartbeat every 10 seconds
if no_update_count % 10 == 0:
progress = TrainingProgressResponse(
step=current_step,
loss=current_loss,
learning_rate=current_lr,
status_message=f"Training step {current_step} (waiting for next update...)"
heartbeat_payload = build_progress(
current_step,
current_loss,
current_lr,
current_total_steps,
current_epoch,
)
yield f"data: {progress.model_dump_json()}\n\n"
yield f"data: {heartbeat_payload.model_dump_json()}\n\n"
else:
# No steps yet, but training is active
no_update_count += 1
if no_update_count % 5 == 0:
yield f"data: {TrainingProgressResponse(step=0, loss=0.0, learning_rate=0.0, status_message='Preparing training...').model_dump_json()}\n\n"
preparing_payload = build_progress(0, 0.0, 0.0, 0)
yield f"data: {preparing_payload.model_dump_json()}\n\n"
# Timeout check
if no_update_count > max_no_updates:
logger.warning("Progress stream timeout - no updates received")
yield f"data: {TrainingProgressResponse(step=last_step, loss=0.0, learning_rate=0.0, status_message='Progress timeout - training may have stopped').model_dump_json()}\n\n"
timeout_payload = build_progress(last_step, 0.0, 0.0, 0)
yield f"data: {timeout_payload.model_dump_json()}\n\n"
break
await asyncio.sleep(1) # Poll every second
except Exception as e:
logger.error(f"Error in progress stream: {e}", exc_info=True)
yield f"data: {TrainingProgressResponse(step=0, loss=0.0, learning_rate=0.0, status_message=f'Error: {str(e)}').model_dump_json()}\n\n"
error_payload = build_progress(0, 0.0, 0.0, 0)
yield f"data: {error_payload.model_dump_json()}\n\n"
break
# Send final status
final_step = backend.step_history[-1] if backend.step_history else last_step
final_loss = backend.loss_history[-1] if backend.loss_history else 0.0
final_lr = backend.lr_history[-1] if backend.lr_history else 0.0
yield f"data: {TrainingProgressResponse(step=final_step, loss=final_loss, learning_rate=final_lr, status_message='Training completed').model_dump_json()}\n\n"
final_tp = getattr(
getattr(backend, "trainer", None), "training_progress", None
)
final_total_steps = (
getattr(final_tp, "total_steps", final_step) if final_tp else final_step
)
final_epoch = getattr(final_tp, "epoch", None) if final_tp else None
final_payload = build_progress(
final_step,
final_loss,
final_lr,
final_total_steps,
final_epoch,
)
yield f"data: {final_payload.model_dump_json()}\n\n"
return StreamingResponse(
event_generator(),

View file

@ -37,7 +37,8 @@ from .chat_templates import (
)
from .vlm_processing import generate_smart_vlm_instruction
from .data_collators import DeepSeekOCRDataCollator, VLMDataCollator
from .model_mappings import TEMPLATE_TO_MODEL_MAPPER, RESPONSE_MARKERS
from .model_mappings import TEMPLATE_TO_MODEL_MAPPER
# , RESPONSE_MARKERS
def check_dataset_format(dataset, is_vlm: bool = False) -> dict: