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:
1 parent
51e74a253e
commit
d170cff0e4
14 files changed
+166
-32
No files matched your search
@@ -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
@@ -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))
|
||||
@@ -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
@@ -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"]
|
||||
|
||||
Reference in new issue
Block a user