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