Files
chicken-counting-sukawarna-det/src/chicken_counter/engine_utils.py
T

240 lines
8.4 KiB
Python

"""TensorRT engine compatibility verification, auto-recompilation, and config updates."""
from __future__ import annotations
import re
import shutil
from pathlib import Path
from typing import Any
import numpy as np
PROJECT_ROOT = Path(__file__).resolve().parents[2]
def verify_engine_compatibility(
model_path: str | Path,
*,
device: str | int | None = "0",
imgsz: int = 640,
) -> tuple[bool, str | None]:
"""Verify if a TensorRT .engine file can be loaded and executed on the current system/GPU.
Returns (True, None) if compatible, or (False, error_reason) if incompatible.
"""
model_path = Path(model_path)
if not model_path.exists():
return False, f"Model file not found: {model_path}"
if model_path.suffix.lower() != ".engine":
return True, None
try:
from ultralytics import YOLO
target_device = str(device) if device is not None else "0"
model = YOLO(str(model_path), task="detect")
dummy_frame = np.zeros((imgsz, imgsz, 3), dtype=np.uint8)
model.predict(dummy_frame, device=target_device, verbose=False)
return True, None
except Exception as exc:
return False, str(exc)
def find_matching_pt_model(
engine_path: str | Path,
search_dirs: list[str | Path] | None = None,
) -> Path | None:
"""Find the matching .pt weights file for a given .engine file."""
engine_path = Path(engine_path)
# 1. Look in the same directory and standard models/ directories
default_dirs = [
engine_path.parent,
PROJECT_ROOT / "models",
PROJECT_ROOT,
]
dirs_to_search = [Path(d).resolve() for d in (search_dirs or default_dirs) if Path(d).exists()]
# 2. Check direct stem match (e.g. model.engine -> model.pt)
direct_pt = engine_path.with_suffix(".pt")
if direct_pt.exists():
return direct_pt.resolve()
for directory in dirs_to_search:
direct = directory / f"{engine_path.stem}.pt"
if direct.exists():
return direct.resolve()
# 3. Check stripped prefix match (e.g. NUC5070_model.engine or jetson_model.engine -> model.pt)
cleaned_stem = re.sub(r"^(NUC\w*|jetson\w*|orin\w*|xavier\w*|nano\w*|arm\w*|x86\w*|gpu\w*)_", "", engine_path.stem, flags=re.IGNORECASE)
for directory in dirs_to_search:
candidate = directory / f"{cleaned_stem}.pt"
if candidate.exists():
return candidate.resolve()
# 4. Search all .pt files in search dirs and find closest substring/stem match
all_pts: list[Path] = []
for directory in dirs_to_search:
all_pts.extend(directory.glob("*.pt"))
if not all_pts:
return None
# Try matching chicken detection models specifically
for pt in all_pts:
if "seg" not in pt.stem.lower() and "pose" not in pt.stem.lower() and "chicken" in pt.stem.lower():
if cleaned_stem.lower() in pt.stem.lower() or pt.stem.lower() in cleaned_stem.lower():
return pt.resolve()
# Fallback to any chicken detection .pt model
for pt in all_pts:
if "seg" not in pt.stem.lower() and "pose" not in pt.stem.lower() and "chicken" in pt.stem.lower():
return pt.resolve()
return all_pts[0].resolve() if all_pts else None
def recompile_engine_from_pt(
pt_path: str | Path,
*,
imgsz: int = 640,
device: str | int | None = "0",
half: bool = True,
workspace: int = 4,
verbose: bool = True,
) -> Path:
"""Compile a new TensorRT .engine from a .pt file on the current machine."""
pt_path = Path(pt_path).resolve()
if not pt_path.exists():
raise FileNotFoundError(f"Source PyTorch model not found: {pt_path}")
from ultralytics import YOLO
target_device = str(device) if device is not None else "0"
print(f"[engine_utils] ⚙️ Compiling TensorRT .engine from: {pt_path.name} (device={target_device}, imgsz={imgsz}, half={half})...")
model = YOLO(str(pt_path), task="detect")
exported_engine = model.export(
format="engine",
imgsz=imgsz,
half=half,
workspace=workspace,
device=target_device,
verbose=verbose,
)
exported_path = Path(exported_engine).resolve()
print(f"[engine_utils] ✅ Successfully compiled TensorRT engine: {exported_path}")
return exported_path
def update_config_yaml_model_path(
config_file_path: str | Path,
new_model_path: str | Path,
) -> bool:
"""Update defaults.detection.model_path in a YAML configuration file while preserving formatting."""
config_path = Path(config_file_path).resolve()
if not config_path.exists():
return False
# Format model path relative to project root if applicable
new_model_path = Path(new_model_path).resolve()
try:
rel_path = new_model_path.relative_to(PROJECT_ROOT)
formatted_path = str(rel_path)
except ValueError:
formatted_path = str(new_model_path)
content = config_path.read_text(encoding="utf-8")
# Match 'model_path: <something>' under detection section
pattern = r"^([ \t]*model_path:[ \t]*)(?:['\"]?)([^'\"\r\n#]+)(?:['\"]?)([ \t]*(?:#.*)?)$"
def replacer(match: re.Match) -> str:
prefix = match.group(1)
comment = match.group(3) or ""
if comment and not comment.startswith(" "):
comment = f" {comment.lstrip()}"
if not comment.startswith(" "):
comment = f" {comment}"
return f"{prefix}{formatted_path}{comment}"
new_content, count = re.subn(pattern, replacer, content, count=1, flags=re.MULTILINE)
if count > 0:
config_path.write_text(new_content, encoding="utf-8")
print(f"[engine_utils] 💾 Updated YAML config '{config_path.name}' -> model_path: {formatted_path}")
return True
# If model_path was not found in this file, check if it extends a base config
extends_match = re.search(r"^\s*(?:extends|base_config):\s*['\"]?([^'\"\s#]+)['\"]?", content, re.MULTILINE)
if extends_match:
base_rel = extends_match.group(1).strip()
base_path = (config_path.parent / base_rel).resolve()
if not base_path.exists():
base_path = (PROJECT_ROOT / base_rel).resolve()
if base_path.exists():
return update_config_yaml_model_path(base_path, new_model_path)
return False
def ensure_compatible_model(
model_path: str | Path,
*,
device: str | int | None = "0",
imgsz: int = 640,
config_file_path: str | Path | None = None,
) -> str:
"""Ensure the model at model_path is compatible with the current hardware.
If an incompatible .engine is detected:
1. Automatically locates the matching .pt file.
2. Recompiles a new .engine optimized for this system.
3. Updates config_file_path (e.g. cycle7_batch_optimized.yaml) with the new engine path.
4. Returns the path to the compatible model.
"""
model_path_obj = Path(model_path)
if model_path_obj.suffix.lower() != ".engine":
return str(model_path)
is_compatible, error_msg = verify_engine_compatibility(
model_path_obj,
device=device,
imgsz=imgsz,
)
if is_compatible:
return str(model_path_obj)
print(f"\n[engine_utils] ⚠️ TensorRT engine '{model_path_obj.name}' is incompatible with this system/GPU.")
print(f"[engine_utils] Reason: {error_msg}")
print("[engine_utils] 🔄 Auto-recompilation triggered: Searching for matching .pt model...")
pt_model = find_matching_pt_model(model_path_obj)
if pt_model is None:
raise RuntimeError(
f"TensorRT engine '{model_path}' is incompatible with this hardware, "
f"and no matching .pt source model was found in {PROJECT_ROOT / 'models'} to recompile from."
)
print(f"[engine_utils] 📦 Found source PyTorch model: {pt_model.name}")
try:
new_engine_path = recompile_engine_from_pt(
pt_model,
imgsz=imgsz,
device=device,
half=True,
workspace=4,
verbose=True,
)
except Exception as export_err:
print(f"[engine_utils] ❌ Engine recompilation failed: {export_err}")
print(f"[engine_utils] ⚠️ Falling back to PyTorch .pt model: {pt_model}")
return str(pt_model)
# If config file is specified, update it
if config_file_path:
update_config_yaml_model_path(config_file_path, new_engine_path)
return str(new_engine_path)