product-detection/3_inference.py

465 lines
19 KiB
Python
Raw Normal View History

2026-05-12 11:16:56 +02:00
"""
Script 3 — End-to-end Inference (Detect → Classify)
=====================================================
1. Runs your YOLOv8n detector on each image / video frame.
2. Crops every detection.
3. Feeds each crop through the classification head.
4. Draws annotated bounding boxes with class names & confidence.
Supported sources
-----------------
• A single image (--source image.jpg)
• A folder of images (--source /path/to/imgs)
• A video file (--source video.mp4)
• A webcam (--source 0)
Usage
-----
python 3_inference.py \
--detector_weights runs/yolo/best.pt \
--classifier_weights runs/classify/best.pt \
--source /path/to/images_or_video \
--det_conf 0.3 \
--cls_conf 0.5 \
--img_size 640 \
--padding 8 \
--output_dir inference_results \
--show \
--save_crops
"""
import argparse
import json
import time
from pathlib import Path
import cv2
import numpy as np
import torch
import torch.nn.functional as F
from torchvision import models, transforms
from ultralytics import YOLO
from tqdm import tqdm
# ─────────────────────────── args ────────────────────────────────────────────
def parse_args():
p = argparse.ArgumentParser(description="YOLOv8 detect + classify inference")
p.add_argument("--detector_weights", default="detector/best.pt", help="YOLOv8n .pt (fine-tuned)")
2026-05-14 15:48:37 +02:00
p.add_argument("--classifier_weights", default="runs/classify/best.pt", help="Classifier best.pt from script 2")
2026-05-12 11:16:56 +02:00
p.add_argument("--source", required=True, help="Image / folder / video / webcam index")
p.add_argument("--det_conf", type=float, default=0.30, help="YOLO detection confidence")
p.add_argument("--det_iou", type=float, default=0.45, help="YOLO NMS IoU")
p.add_argument("--cls_conf", type=float, default=0.60, help="Min classifier confidence to label")
p.add_argument("--img_size", type=int, default=640, help="YOLO input size")
p.add_argument("--padding", type=int, default=8, help="Extra pixels around each crop")
p.add_argument("--cls_img_size",type=int, default=224, help="Classifier input size")
p.add_argument("--output_dir", default="inference_results")
p.add_argument("--show", action="store_true", help="Display results with cv2.imshow")
p.add_argument("--save_crops", action="store_true", help="Save individual crop images")
p.add_argument("--no_save", action="store_true", help="Do not save annotated images/video")
p.add_argument("--batch_cls", type=int, default=16, help="Classifier batch size per frame")
p.add_argument("--device", default="", help="cuda / cpu / mps (auto-detect if empty)")
return p.parse_args()
# ─────────────────────────── classifier wrapper ──────────────────────────────
# ── Device-agnostic OOM detection ────────────────────────────────────────────
#
# Each backend raises a different exception type for out-of-memory:
#
# CUDA → torch.cuda.OutOfMemoryError (subclass of RuntimeError)
# MPS → RuntimeError whose message contains "out of memory"
# CPU → MemoryError (Python built-in, raised by the OS allocator)
#
# We centralise the check here so nothing else needs to know the device type.
def _is_oom(exc: BaseException) -> bool:
"""Return True if *exc* represents an out-of-memory condition."""
if isinstance(exc, MemoryError): # CPU / OS
return True
if torch.cuda.is_available() and isinstance(exc, torch.cuda.OutOfMemoryError):
return True # CUDA
if isinstance(exc, RuntimeError): # MPS + fallback CUDA
msg = str(exc).lower()
return "out of memory" in msg or "memory" in msg and "alloc" in msg
return False
def _free_device_cache(device: torch.device) -> None:
"""Release any cached memory held by the current device allocator."""
if device.type == "cuda":
torch.cuda.empty_cache()
elif device.type == "mps":
# torch.mps.empty_cache() was added in PyTorch 2.1 — guard for older versions
if hasattr(torch.mps, "empty_cache"):
torch.mps.empty_cache()
# CPU has no cache to release
def _sync_device(device: torch.device) -> None:
"""Wait for all pending ops on *device* so OOM surfaces immediately."""
if device.type == "cuda":
torch.cuda.synchronize()
elif device.type == "mps":
if hasattr(torch.mps, "synchronize"):
torch.mps.synchronize()
# ─────────────────────────────────────────────────────────────────────────────
class ProductClassifier:
"""Thin wrapper around the saved checkpoint from script 2."""
# Supported architectures — mirrors script 2
_BUILDERS = {
"efficientnet_b0": (models.efficientnet_b0, lambda m, n: setattr(m.classifier, '1', torch.nn.Linear(m.classifier[1].in_features if hasattr(m.classifier[1], 'in_features') else 1280, n))),
"efficientnet_b2": (models.efficientnet_b2, lambda m, n: setattr(m.classifier, '1', torch.nn.Linear(m.classifier[1].in_features if hasattr(m.classifier[1], 'in_features') else 1408, n))),
"mobilenet_v3_small": (models.mobilenet_v3_small, lambda m, n: setattr(m.classifier, '3', torch.nn.Linear(m.classifier[3].in_features if hasattr(m.classifier[3], 'in_features') else 576, n))),
"resnet50": (models.resnet50, lambda m, n: setattr(m, 'fc', torch.nn.Linear(m.fc.in_features, n))),
}
def __init__(self, weights_path: str, device: torch.device, img_size: int = 224):
ckpt = torch.load(weights_path, map_location=device)
self.class_names: list = ckpt["class_names"]
num_classes = len(self.class_names)
model_name = ckpt.get("args", {}).get("model", "efficientnet_b0")
builder_fn, head_fn = self._BUILDERS[model_name]
self.model = builder_fn(weights=None)
head_fn(self.model, num_classes)
self.model.load_state_dict(ckpt["model_state"])
self.model.to(device).eval()
self.device = device
mean = [0.485, 0.456, 0.406]
std = [0.229, 0.224, 0.225]
self.transform = transforms.Compose([
transforms.ToPILImage(),
transforms.Resize(int(img_size * 1.15)),
transforms.CenterCrop(img_size),
transforms.ToTensor(),
transforms.Normalize(mean, std),
])
# Probe VRAM once at startup — stored and reused every frame.
# Can be overridden at runtime if OOM occurs (see predict_batch).
self._batch_size: int = self.calibrate_batch_size(img_size)
def calibrate_batch_size(self, img_size: int, start: int = 2, max_size: int = 512) -> int:
"""
Probe the device at startup to find the largest batch that fits in memory.
Works on CUDA, MPS, and CPU:
• CUDA / MPS — doubles the batch size until OOM, steps back with margin.
• CPU — RAM is virtually unlimited; returns a sensible fixed cap.
Returns the largest safe batch size, always ≥ 1.
"""
if self.device.type == "cpu":
print(" [batch-probe] CPU detected — using fixed batch size 64")
return 64
print(f" [batch-probe] Probing {self.device.type.upper()} memory "
f"for optimal batch size (img={img_size}px) …", flush=True)
dummy = torch.zeros(1, 3, img_size, img_size, device=self.device)
safe_size = 1
probe_size = start
with torch.no_grad():
while probe_size <= max_size:
try:
batch = dummy.expand(probe_size, -1, -1, -1)
_ = self.model(batch)
# Flush pending async ops so OOM surfaces here, not later
_sync_device(self.device)
safe_size = probe_size
probe_size = probe_size * 2
except Exception as exc:
if _is_oom(exc):
break # safe_size is our answer
raise # unexpected error — propagate
finally:
_free_device_cache(self.device)
# 20 % headroom margin: real crops vary in size unlike the uniform dummy
safe_size = max(1, int(safe_size * 0.8))
print(f" [batch-probe] Settled on batch_size = {safe_size} "
f"(probe ceiling was {probe_size // 2})")
return safe_size
@torch.no_grad()
def predict_batch(self, crops_bgr: list, batch_size: int | None = None) -> list[tuple[str, float]]:
"""
Classifies an arbitrary list of crops in memory-safe chunks.
crops_bgr : list of H×W×3 uint8 BGR numpy arrays
batch_size : chunk size to use. Pass None to use self._batch_size
(set by calibrate_batch_size at startup).
Returns : list of (class_name, confidence) — same order as input
"""
if not crops_bgr:
return []
# Resolve which batch size to use
chunk_size = batch_size if batch_size is not None else self._batch_size
# Pre-process all crops on CPU (fast; no GPU memory involved yet)
tensors = [
self.transform(cv2.cvtColor(crop, cv2.COLOR_BGR2RGB))
for crop in crops_bgr
]
results: list[tuple[str, float]] = []
i = 0
while i < len(tensors):
chunk = tensors[i : i + chunk_size]
batch = torch.stack(chunk).to(self.device)
try:
logits = self.model(batch)
probs = F.softmax(logits, dim=1)
confs, idxs = probs.max(dim=1)
for idx, conf in zip(idxs, confs):
results.append((self.class_names[idx.item()], conf.item()))
i += chunk_size # advance only on success
except Exception as exc:
if not _is_oom(exc):
raise # unexpected error — don't swallow it
# ── Runtime OOM safety net ────────────────────────────────
# The probe uses a uniform dummy; real crops vary and can
# occasionally exceed the probed limit. Halve and retry
# WITHOUT advancing i so the same chunk is retried.
_free_device_cache(self.device)
new_size = max(1, chunk_size // 2)
print(f" [batch] Runtime OOM on {self.device.type.upper()} "
f"— halving chunk {chunk_size} → {new_size}")
chunk_size = new_size
self._batch_size = new_size # persist for future frames
finally:
del batch # free GPU memory immediately after each chunk
return results
# ─────────────────────────── drawing ─────────────────────────────────────────
# One colour per class, generated on first use
_COLOUR_CACHE: dict[str, tuple] = {}
def class_colour(name: str) -> tuple:
if name not in _COLOUR_CACHE:
# OpenCV HSV: hue is 0–179 (NOT 0–360) — values above 179 overflow uint8
h = hash(name) % 180
col = cv2.cvtColor(
np.array([[[h, 220, 200]]], dtype=np.uint8), cv2.COLOR_HSV2BGR
)[0][0]
_COLOUR_CACHE[name] = tuple(int(c) for c in col)
return _COLOUR_CACHE[name]
def draw_result(frame, box_xyxy, label: str, conf: float, cls_conf_thresh: float):
x1, y1, x2, y2 = map(int, box_xyxy)
colour = class_colour(label)
cv2.rectangle(frame, (x1, y1), (x2, y2), colour, 2)
if conf >= cls_conf_thresh:
text = f"{label} {conf:.2f}"
else:
text = "?"
(tw, th), bl = cv2.getTextSize(text, cv2.FONT_HERSHEY_SIMPLEX, 0.55, 1)
cv2.rectangle(frame, (x1, y1 - th - bl - 4), (x1 + tw + 4, y1), colour, -1)
cv2.putText(frame, text, (x1 + 2, y1 - bl - 2),
cv2.FONT_HERSHEY_SIMPLEX, 0.55, (255, 255, 255), 1, cv2.LINE_AA)
return frame
# ─────────────────────────── source helpers ──────────────────────────────────
IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".bmp", ".webp", ".tiff"}
def is_image(path: Path) -> bool:
return path.suffix.lower() in IMAGE_EXTS
def is_video(path: Path) -> bool:
return path.suffix.lower() in {".mp4", ".avi", ".mov", ".mkv", ".webm"}
def iter_source(source: str):
"""
Yields (frame_bgr, frame_id, source_name) for every frame / image.
"""
try:
cam_idx = int(source) # webcam
cap = cv2.VideoCapture(cam_idx)
fid = 0
while True:
ret, frame = cap.read()
if not ret:
break
yield frame, fid, f"webcam_{cam_idx}"
fid += 1
cap.release()
return
except ValueError:
pass
p = Path(source)
if p.is_file() and is_image(p):
img = cv2.imread(str(p))
if img is not None:
yield img, 0, p.stem
elif p.is_file() and is_video(p):
cap = cv2.VideoCapture(str(p))
fid = 0
while True:
ret, frame = cap.read()
if not ret:
break
yield frame, fid, p.stem
fid += 1
cap.release()
elif p.is_dir():
img_paths = sorted(q for q in p.rglob("*") if is_image(q))
for q in tqdm(img_paths, desc="Images"):
img = cv2.imread(str(q))
if img is not None:
yield img, 0, q.stem
else:
raise FileNotFoundError(f"Cannot open source: {source}")
# ─────────────────────────── main ────────────────────────────────────────────
def main():
args = parse_args()
# Device
if args.device:
device = torch.device(args.device)
elif torch.cuda.is_available():
device = torch.device("cuda")
elif torch.backends.mps.is_available():
device = torch.device("mps")
else:
device = torch.device("cpu")
print(f"Device : {device}")
# Load models
print(f"Loading detector : {args.detector_weights}")
detector = YOLO(args.detector_weights)
print(f"Loading classifier: {args.classifier_weights}")
classifier = ProductClassifier(args.classifier_weights, device, args.cls_img_size)
print(f" Classes ({len(classifier.class_names)}): {classifier.class_names}")
# Output directory
out_dir = Path(args.output_dir)
if not args.no_save:
out_dir.mkdir(parents=True, exist_ok=True)
crops_dir = out_dir / "crops" if args.save_crops else None
if crops_dir:
crops_dir.mkdir(parents=True, exist_ok=True)
# Stats
stats = {"frames": 0, "detections": 0, "classified": 0, "fps": []}
for frame, frame_id, stem in iter_source(args.source):
t0 = time.perf_counter()
# ── Detection ──────────────────────────────────────────────────────
results = detector(
frame,
conf=args.det_conf,
iou=args.det_iou,
imgsz=args.img_size,
verbose=False,
)[0]
h, w = frame.shape[:2]
crops, boxes = [], []
for box in results.boxes:
x1, y1, x2, y2 = map(int, box.xyxy[0].tolist())
x1 = max(0, x1 - args.padding); y1 = max(0, y1 - args.padding)
x2 = min(w, x2 + args.padding); y2 = min(h, y2 + args.padding)
crop = frame[y1:y2, x1:x2]
if crop.size == 0:
continue
crops.append(crop)
boxes.append((x1, y1, x2, y2))
# ── Classification (dynamic chunked batches) ───────────────────────
# Batch size is auto-calibrated to VRAM at startup via probe,
# and auto-reduces at runtime if OOM occurs (e.g. unusually dense frame).
predictions = classifier.predict_batch(crops)
# ── Draw & annotate ───────────────────────────────────────────────
annotated = frame.copy()
for (x1, y1, x2, y2), (cls_name, cls_conf) in zip(boxes, predictions):
draw_result(annotated, (x1, y1, x2, y2), cls_name, cls_conf, args.cls_conf)
if cls_conf >= args.cls_conf:
stats["classified"] += 1
# FPS overlay
elapsed = time.perf_counter() - t0
fps = 1.0 / elapsed if elapsed > 0 else 0
stats["fps"].append(fps)
cv2.putText(annotated, f"FPS: {fps:.1f} Det: {len(boxes)}",
(10, 28), cv2.FONT_HERSHEY_SIMPLEX, 0.8, (0, 255, 0), 2)
# ── Save / display ────────────────────────────────────────────────
if not args.no_save:
out_path = out_dir / f"{stem}_frame{frame_id:06d}.jpg"
cv2.imwrite(str(out_path), annotated)
if args.save_crops and crops_dir:
for i, (crop, (cls_name, cls_conf)) in enumerate(zip(crops, predictions)):
tag = cls_name if cls_conf >= args.cls_conf else "uncertain"
cv2.imwrite(str(crops_dir / f"{stem}_f{frame_id:06d}_c{i:04d}_{tag}.jpg"), crop)
if args.show:
cv2.imshow("Detect + Classify", annotated)
key = cv2.waitKey(1)
if key in (ord("q"), 27):
break
stats["frames"] += 1
stats["detections"] += len(boxes)
cv2.destroyAllWindows()
# ── Summary ───────────────────────────────────────────────────────────
avg_fps = np.mean(stats["fps"]) if stats["fps"] else 0
print("\n" + "═" * 55)
print(" Inference summary")
print(f" Frames processed : {stats['frames']}")
print(f" Total detections : {stats['detections']}")
print(f" Classified (≥{args.cls_conf:.0%}): {stats['classified']}")
print(f" Avg FPS : {avg_fps:.1f}")
if not args.no_save:
print(f" Results saved to : {out_dir.resolve()}")
print("═" * 55)
if __name__ == "__main__":
main()