"""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: ' 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)