forked from zakaria/chicken-counting-sukawarna-det
240 lines
8.4 KiB
Python
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)
|