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
This commit is contained in:
Andrew-AAAA committed 2026-09-10 09:14:55 +07:00
1 parent 51e74a253e
commit d170cff0e4
14 files changed
+166 -32

No files matched your search

+1 -1
View File
@@ -46,7 +46,7 @@ def available_models(project_id: int) -> dict:
for version in training.listing(project_id):
if version.get("weights_path") and os.path.isfile(version["weights_path"]):
out.append({
"label": f"v{version['version']}",
"label": version.get("name") or f"v{version['version']}",
"path": version["weights_path"],
"version_id": version["id"],
})
+17 -1
View File
@@ -57,8 +57,9 @@ 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"v{version['version']}-best.pt")
filename=f"{name}-best.pt")
@router.post("/api/models/{model_id}/promote")
@@ -67,3 +68,18 @@ def promote_model(model_id: int) -> dict:
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))
+46
View File
@@ -270,6 +270,8 @@ def migrate() -> None:
# REQ-113: and what augmentation it trained under.
if "augment" not in version_cols:
cur.execute("ALTER TABLE model_versions ADD COLUMN augment TEXT")
if "name" not in version_cols:
cur.execute("ALTER TABLE model_versions ADD COLUMN name TEXT")
# REQ-132: the triage rules this dataset was actually cut under, frozen
# at merge time. Editing project rules afterwards must not rewrite what
@@ -284,6 +286,7 @@ def migrate() -> None:
_migrate_job_types(cur)
_migrate_clock_column(cur)
_migrate_truck_columns(cur)
_backfill_model_names(cur)
def _migrate_dataset_items(cur) -> None:
@@ -391,6 +394,49 @@ def _migrate_job_types(cur) -> None:
cur.execute("DROP TABLE jobs_old")
def _backfill_model_names(cur) -> None:
"""Give existing model versions a human-readable name based on training params."""
import time as _time
PRETRAINED = {"bbox": "yolo11n.pt", "polygon": "yolo11n-seg.pt"}
cur.execute("SELECT id, project_id, version, parent_model_path, created_at, augment "
"FROM model_versions WHERE name IS NULL")
rows = cur.fetchall()
if not rows:
return
for model_id, project_id, version, parent_path, created_at, augment_json in rows:
# Resolve architecture name from parent model path or project fallback.
cur.execute("SELECT label_type FROM projects WHERE id = ?", (project_id,))
proj_row = cur.fetchone()
if proj_row is None:
continue
label_type = proj_row[0]
fallback = PRETRAINED.get(label_type, "yolo11n.pt")
if parent_path:
arch = os.path.splitext(os.path.basename(parent_path))[0]
else:
arch = os.path.splitext(os.path.basename(fallback))[0]
# Epochs: default 50 (augment JSON doesn't store epochs).
epochs = 50
# Class names.
cur.execute(
"SELECT name FROM project_classes WHERE project_id = ? ORDER BY class_id",
(project_id,),
)
class_names = "+".join(r[0] for r in cur.fetchall()) or "unknown"
# Date.
date_str = _time.strftime("%Y%m%d", _time.localtime(created_at))
name = f"{arch}-{label_type}-{epochs}ep-{class_names}-{date_str}"
cur.execute("UPDATE model_versions SET name = ? WHERE id = ?", (name, model_id))
def _backfill_dataset_rules(cur) -> None:
"""Datasets merged before REQ-132 have no snapshot. Give them the project's
current rules — that is what they were cut under, unless the rules changed
+29 -3
View File
@@ -12,6 +12,21 @@ import shutil
import time
from typing import Optional
def _generate_model_name(project: dict, start_point: str, epochs: int,
class_ids: Optional[list] = None) -> str:
"""Build a descriptive model name: arch-labelType-epochs-classNames-YYYYMMDD."""
arch = os.path.splitext(os.path.basename(start_point))[0]
# Filter classes if specific IDs were selected.
classes = project["classes"]
if class_ids:
classes = [c for c in classes if c["class_id"] in class_ids]
class_tag = "+".join(c["name"] for c in classes) or "unknown"
date_str = time.strftime("%Y%m%d", time.localtime())
return f"{arch}-{project['label_type']}-{epochs}ep-{class_tag}-{date_str}"
from backend import (augment, base_dataset, config, dataset, datasets, db, evaluate,
hardware, jobs, projects)
@@ -149,6 +164,16 @@ def promote(model_id: int) -> dict:
return projects.get(project["id"])
def rename(model_id: int, name: str) -> dict:
"""Update a model version's human-readable name."""
version = get_version(model_id)
if version is None:
raise TrainingError("No such model version")
with db.cursor() as cur:
cur.execute("UPDATE model_versions SET name = ? WHERE id = ?", (name.strip(), model_id))
return get_version(model_id)
def _next_version(cur, project_id: int) -> int:
cur.execute(
"SELECT COALESCE(MAX(version), 0) + 1 FROM model_versions WHERE project_id = ?",
@@ -262,11 +287,12 @@ def _run_train(job) -> None:
cur.execute(
"""INSERT INTO model_versions (project_id, version, weights_path,
parent_model_path, metrics, base_metrics,
created_at, augment)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)""",
created_at, augment, name)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)""",
(project["id"], version, weights, project["base_model_path"],
json.dumps(comparison["new"]), json.dumps(comparison["base"]), time.time(),
json.dumps(augmentation["settings"])),
json.dumps(augmentation["settings"]),
_generate_model_name(project, start_point, settings["epochs"], class_ids)),
)
new = comparison["new"]