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:
1 parent
8f41c6c85a
commit
dee58e4ae5
10 files changed
+93
-32
No files matched your search
@@ -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")
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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)),
|
||||
|
||||
Reference in new issue
Block a user