- config.py: resolve_data_path (legacy abs + rel) + rel_data_path - all file-opening reads wrapped: preview, autolabel, training, model download, live count, projects.get; training_start_point hack replaced - new writes store paths relative to data/ - legacy stale rows (/home/asus/reTraining/...) resolve without migration - requirements: REQ-187 added; REQ-188 (per-class max box) + REQ-186 copy-line amendment drafted for the next task
87 lines
2.7 KiB
Python
87 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 config, 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)
|
|
weights = config.resolve_data_path(version["weights_path"]) if version else None
|
|
if version is None or not os.path.isfile(weights):
|
|
raise HTTPException(404, "No weights for that version")
|
|
name = version.get('name') or f"v{version['version']}"
|
|
return FileResponse(weights, 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))
|