Files
reTraining/backend/api/models.py
T
asus dee58e4ae5 fix: resolve model weight paths against DATA_DIR (REQ-187)
- 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
2026-10-02 16:59:25 +07:00

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))