Files
reTraining/backend/api/models.py
T
Andrew-AAAA d170cff0e4 feat: add descriptive model naming and inline rename
- Auto-generate model names: {arch}-{labelType}-{epochs}ep-{classNames}-{YYYYMMDD}
- Add PATCH /api/models/{id}/rename endpoint
- Inline rename UI on Models & Training page
- Download filename uses model name instead of v{N}
- DB migration: add name column to model_versions
- Update all docs to reflect new naming convention
2026-09-10 09:14:55 +07:00

86 lines
2.7 KiB
Python

"""Training and model-version routes (REQ-060…065)."""
import os
from typing import Optional, Union
from fastapi import APIRouter, HTTPException
from fastapi.responses import FileResponse
from pydantic import BaseModel
from backend import hardware, training
from backend.api.common import project_or_404
router = APIRouter(tags=["models"])
class TrainRequest(BaseModel):
epochs: int = 50
batch: Optional[int] = None
imgsz: Optional[int] = None
device: Optional[Union[int, str]] = None
batch_ids: Optional[list] = None
class_ids: Optional[list] = None
dataset_ids: Optional[list] = None
base_dataset_ids: Optional[list] = None
@router.get("/api/hardware")
def read_hardware() -> dict:
return hardware.defaults()
@router.post("/api/projects/{project_id}/train")
def start_training(project_id: int, request: TrainRequest) -> dict:
project_or_404(project_id)
try:
return training.start(
project_id, request.epochs,
{"batch": request.batch, "imgsz": request.imgsz, "device": request.device},
batch_ids=request.batch_ids,
class_ids=request.class_ids,
dataset_ids=request.dataset_ids,
base_dataset_ids=request.base_dataset_ids,
)
except training.TrainingError as exc:
raise HTTPException(400, str(exc))
@router.get("/api/projects/{project_id}/models")
def list_models(project_id: int) -> dict:
project = project_or_404(project_id)
return {"models": training.listing(project_id),
"base_model_kind": project["base_model_kind"]}
@router.get("/api/models/{model_id}/weights")
def download_weights(model_id: int):
version = training.get_version(model_id)
if version is None or not os.path.isfile(version["weights_path"]):
raise HTTPException(404, "No weights for that version")
name = version.get('name') or f"v{version['version']}"
return FileResponse(version["weights_path"], media_type="application/octet-stream",
filename=f"{name}-best.pt")
@router.post("/api/models/{model_id}/promote")
def promote_model(model_id: int) -> dict:
try:
return training.promote(model_id)
except training.TrainingError as exc:
raise HTTPException(400, str(exc))
class RenameRequest(BaseModel):
name: str
@router.patch("/api/models/{model_id}/rename")
def rename_model(model_id: int, request: RenameRequest) -> dict:
try:
version = training.rename(model_id, request.name)
if version is None:
raise HTTPException(404, "No such model version")
return version
except training.TrainingError as exc:
raise HTTPException(400, str(exc))