66 lines
2.0 KiB
Python
66 lines
2.0 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
|
|
|
|
|
|
@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,
|
|
)
|
|
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")
|
|
return FileResponse(version["weights_path"], media_type="application/octet-stream",
|
|
filename=f"v{version['version']}-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))
|