product-detection/3_inference.py

465 lines
19 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
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)")
p.add_argument("--classifier_weights", default="runs/classify/best.pt", help="Classifier best.pt from script 2")
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 0179 (NOT 0360) — 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()