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
This commit is contained in:
asus committed 2026-10-02 16:59:25 +07:00
1 parent 8f41c6c85a
commit dee58e4ae5
10 files changed
+93 -32

No files matched your search

+4 -3
View File
@@ -8,7 +8,7 @@ from fastapi.responses import StreamingResponse
from pydantic import BaseModel
from typing import Optional
from backend import library, live_count, training
from backend import config, library, live_count, training
from backend.api.common import project_or_404
router = APIRouter(tags=["live-count"])
@@ -45,10 +45,11 @@ def available_models(project_id: int) -> dict:
project = project_or_404(project_id)
out = []
for version in training.listing(project_id):
if version.get("weights_path") and os.path.isfile(version["weights_path"]):
path = config.resolve_data_path(version["weights_path"])
if version.get("weights_path") and os.path.isfile(path):
out.append({
"label": version.get("name") or f"v{version['version']}",
"path": version["weights_path"],
"path": path,
"version_id": version["id"],
})
base = project.get("base_model_path")
+4 -3
View File
@@ -7,7 +7,7 @@ from fastapi import APIRouter, HTTPException
from fastapi.responses import FileResponse
from pydantic import BaseModel
from backend import hardware, training
from backend import config, hardware, training
from backend.api.common import project_or_404
router = APIRouter(tags=["models"])
@@ -55,10 +55,11 @@ def list_models(project_id: int) -> dict:
@router.get("/api/models/{model_id}/weights")
def download_weights(model_id: int):
version = training.get_version(model_id)
if version is None or not os.path.isfile(version["weights_path"]):
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(version["weights_path"], media_type="application/octet-stream",
return FileResponse(weights, media_type="application/octet-stream",
filename=f"{name}-best.pt")
+5 -3
View File
@@ -10,7 +10,7 @@ import os
from typing import List, Optional
from PIL import Image
from backend import batches, db, jobs, labeling, projects, review
from backend import batches, config, db, jobs, labeling, projects, review
from backend.batches import BatchError
DEFAULT_THRESHOLD = 0.35
@@ -158,8 +158,10 @@ def _run_autolabel(job) -> None:
with db.cursor() as cur:
cur.execute("SELECT weights_path FROM model_versions WHERE project_id = ? ORDER BY version DESC LIMIT 1", (project["id"],))
row = cur.fetchone()
if row and os.path.isfile(row[0]):
m_path = row[0]
if row:
resolved = config.resolve_data_path(row[0])
if os.path.isfile(resolved):
m_path = resolved
job.log(f"Loading Base Model: {os.path.basename(m_path)}...")
yolo_model = YOLO(m_path)
+28
View File
@@ -18,6 +18,34 @@ DATA_DIR = os.path.abspath(os.environ.get("APP_DATA_DIR", os.path.join(REPO_ROOT
PROJECTS_DIR = os.path.join(DATA_DIR, "projects")
DB_PATH = os.path.join(DATA_DIR, "app.db")
def resolve_data_path(path):
"""Resolve a DB-stored weight path against today's DATA_DIR (REQ-187)."""
if not path:
return path
if os.path.isfile(path):
return path
if not os.path.isabs(path):
candidate = os.path.join(DATA_DIR, path)
return candidate if os.path.isfile(candidate) else path
# legacy absolute path from before the data dir moved: .../data/<rel>
marker = os.sep + "data" + os.sep
idx = path.find(marker)
if idx != -1:
candidate = os.path.join(DATA_DIR, path[idx + len(marker):])
if os.path.isfile(candidate):
return candidate
return path
def rel_data_path(path):
"""Store paths under DATA_DIR relative to it (REQ-187); others unchanged."""
try:
rel = os.path.relpath(path, DATA_DIR)
except ValueError:
return path
return path if rel.startswith(os.pardir) else rel
# Where the video archive is mounted. Projects store a path relative to nothing —
# they store an absolute one — but this is the default the UI starts browsing from.
VIDEO_ROOT = os.path.abspath(os.environ.get("VIDEO_ARCHIVE", os.path.join(DATA_DIR, "archive")))
+5 -3
View File
@@ -11,7 +11,7 @@ whatever happens to sit at those coordinates there.
import os
from typing import List, Optional
from backend import batches, db, labeling, projects, review
from backend import batches, config, db, labeling, projects, review
from backend.autolabel import DEFAULT_IOU, DEFAULT_THRESHOLD, _geometries, _parse_class_params
def preview_frame(
@@ -64,8 +64,10 @@ def preview_frame(
with db.cursor() as cur:
cur.execute("SELECT weights_path FROM model_versions WHERE project_id = ? ORDER BY version DESC LIMIT 1", (project["id"],))
row = cur.fetchone()
if row and os.path.isfile(row[0]):
m_path = row[0]
if row:
resolved = config.resolve_data_path(row[0])
if os.path.isfile(resolved):
m_path = resolved
yolo_model = YOLO(m_path)
name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
+5 -10
View File
@@ -112,8 +112,8 @@ def create(name: str, label_type: str, video_root: str, classes: Optional[List[d
"""INSERT INTO projects (slug, name, label_type, base_model_path,
base_model_kind, video_root, val_every, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)""",
(slug, name.strip(), label_type, stored_model, kind, video_root,
max(0, val_every), time.time()),
(slug, name.strip(), label_type, config.rel_data_path(stored_model), kind,
video_root, max(0, val_every), time.time()),
)
project_id = cur.lastrowid
_write_classes(cur, project_id, cleaned)
@@ -173,7 +173,7 @@ def _row_to_dict(cur, row) -> dict:
"slug": row["slug"],
"name": row["name"],
"label_type": row["label_type"],
"base_model_path": row["base_model_path"],
"base_model_path": config.resolve_data_path(row["base_model_path"]),
"base_model_kind": row["base_model_kind"],
"secondary_model_path": row["secondary_model_path"] if "secondary_model_path" in row.keys() else None,
"secondary_model_name": row["secondary_model_name"] if "secondary_model_name" in row.keys() else None,
@@ -350,7 +350,7 @@ def set_base_model(project_id: int, weights_path: str) -> dict:
with db.cursor() as cur:
cur.execute(
"UPDATE projects SET base_model_path = ?, base_model_kind = 'uploaded' WHERE id = ?",
(stored, project_id),
(config.rel_data_path(stored), project_id),
)
_write_classes(cur, project_id, _kept_classes(project, names))
return get(project_id)
@@ -391,12 +391,7 @@ def _kept_classes(project: dict, names: List[str]) -> List[dict]:
def training_start_point(project: dict) -> str:
"""The weights a training run should start from (REQ-060, REQ-004)."""
path = project["base_model_path"]
if path and not os.path.isfile(path):
if path.startswith("/data/"):
alt_path = os.path.join(config.DATA_DIR, path[6:])
if os.path.isfile(alt_path):
path = alt_path
path = config.resolve_data_path(project["base_model_path"])
if path and os.path.isfile(path):
return path
return PRETRAINED[project["label_type"]]
+5 -3
View File
@@ -155,11 +155,11 @@ def promote(model_id: int) -> dict:
project = projects.get(version["project_id"])
base_path = os.path.join(config.project_dir(project["slug"]), "base", "model.pt")
os.makedirs(os.path.dirname(base_path), exist_ok=True)
shutil.copyfile(version["weights_path"], base_path)
shutil.copyfile(config.resolve_data_path(version["weights_path"]), base_path)
with db.cursor() as cur:
cur.execute(
"UPDATE projects SET base_model_path = ?, base_model_kind = 'trained' WHERE id = ?",
(base_path, project["id"]),
(config.rel_data_path(base_path), project["id"]),
)
return projects.get(project["id"])
@@ -289,7 +289,9 @@ def _run_train(job) -> None:
parent_model_path, metrics, base_metrics,
created_at, augment, name)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)""",
(project["id"], version, weights, project["base_model_path"],
(project["id"], version, config.rel_data_path(weights),
config.rel_data_path(project["base_model_path"])
if project["base_model_path"] else project["base_model_path"],
json.dumps(comparison["new"]), json.dumps(comparison["base"]), time.time(),
json.dumps(augmentation["settings"]),
_generate_model_name(project, start_point, settings["epochs"], class_ids)),