This commit includes major additions and updates to the frontend and backend architectures, introducing new dataset management, live counting features, batch processing, and triage logic. Includes new UI pages, components, and API routes.
70 lines
2.2 KiB
Python
70 lines
2.2 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")
|
|
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))
|