HaWkEye / appearance_engine.py
HaWkEye Admin
Deploy HaWkEye
97c8cf9
Raw
History Blame Contribute Delete
5.79 kB
"""
Appearance Re-Identification Engine (OSNet-AIN)
================================================
Extracts 512-dimensional appearance embeddings using the OSNet-AIN model
from torchreid, pretrained on Market-1501 + DukeMTMC.
OSNet learns clothing-independent features (body shape, posture, proportions)
at multiple spatial scales via its Omni-Scale architecture.
Falls back to ResNet18 feature extraction if torchreid is not available.
"""
import os
import cv2
import torch
import numpy as np
from PIL import Image
from torchvision import transforms
# Try loading torchreid
TORCHREID_AVAILABLE = False
try:
import torchreid
TORCHREID_AVAILABLE = True
except ImportError:
print("[AppearanceEngine] torchreid not found — will use ResNet18 fallback")
class AppearanceEngine:
"""
Extracts 512-d appearance embeddings for person re-identification.
Primary: OSNet-AIN (torchreid) — state-of-the-art person re-ID
Fallback: ResNet18 (torchvision) — basic appearance features
"""
EMBEDDING_DIM = 512
INPUT_SIZE = (256, 128) # H, W — standard person re-ID input
def __init__(self, device=None):
self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
self.model = None
self.model_name = "none"
# Image preprocessing
self.transform = transforms.Compose([
transforms.Resize(self.INPUT_SIZE),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]),
])
self._load_model()
def _load_model(self):
"""Load OSNet-AIN or fallback to ResNet18."""
if TORCHREID_AVAILABLE:
try:
print("[AppearanceEngine] Loading OSNet-AIN (torchreid)...")
self.model = torchreid.models.build_model(
name="osnet_ain_x1_0",
num_classes=1000, # doesn't matter for feature extraction
pretrained=True,
)
self.model = self.model.to(self.device)
self.model.eval()
self.model_name = "OSNet-AIN"
self.EMBEDDING_DIM = 512
print("[AppearanceEngine] OSNet-AIN loaded successfully")
return
except Exception as e:
print(f"[AppearanceEngine] OSNet-AIN load failed: {e}")
# Fallback: ResNet18
try:
print("[AppearanceEngine] Loading ResNet18 fallback...")
from torchvision import models
self.model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)
self.model.fc = torch.nn.Identity() # remove classification head
self.model = self.model.to(self.device)
self.model.eval()
self.model_name = "ResNet18"
self.EMBEDDING_DIM = 512
print("[AppearanceEngine] ResNet18 fallback loaded")
except Exception as e:
print(f"[AppearanceEngine] All model loads failed: {e}")
@torch.no_grad()
def extract_embedding(self, person_crop_bgr):
"""
Extract appearance embedding from a person crop.
Args:
person_crop_bgr: BGR OpenCV image (person body crop)
Returns:
np.array of shape (512,), L2-normalized. None on failure.
"""
if self.model is None or person_crop_bgr is None:
return None
if person_crop_bgr.size == 0:
return None
try:
rgb = cv2.cvtColor(person_crop_bgr, cv2.COLOR_BGR2RGB)
pil_img = Image.fromarray(rgb)
tensor = self.transform(pil_img).unsqueeze(0).to(self.device)
if self.model_name == "OSNet-AIN":
# torchreid models: use feature extraction mode
features = self.model(tensor)
if isinstance(features, (tuple, list)):
features = features[0]
emb = features.cpu().numpy().flatten()
else:
# ResNet18 fallback
emb = self.model(tensor).cpu().numpy().flatten()
# L2 normalize
norm = np.linalg.norm(emb)
if norm > 0:
emb = emb / norm
# Pad or truncate to target dim
if len(emb) < self.EMBEDDING_DIM:
emb = np.pad(emb, (0, self.EMBEDDING_DIM - len(emb)))
elif len(emb) > self.EMBEDDING_DIM:
emb = emb[:self.EMBEDDING_DIM]
return emb.astype(np.float32)
except Exception as e:
print(f"[AppearanceEngine] Embedding extraction failed: {e}")
return None
def extract_embedding_batch(self, crops_bgr):
"""Extract embeddings for a batch of person crops."""
embeddings = []
for crop in crops_bgr:
emb = self.extract_embedding(crop)
if emb is not None:
embeddings.append(emb)
return embeddings
@staticmethod
def compute_similarity(emb1, emb2):
"""Cosine similarity between two appearance embeddings."""
if emb1 is None or emb2 is None:
return 0.0
n1 = np.linalg.norm(emb1)
n2 = np.linalg.norm(emb2)
if n1 < 1e-8 or n2 < 1e-8:
return 0.0
return float(np.dot(emb1, emb2) / (n1 * n2))
@staticmethod
def average_embeddings(embeddings):
"""Average a list of embeddings into a single template."""
if not embeddings:
return None
stacked = np.stack(embeddings, axis=0)
avg = np.mean(stacked, axis=0)
norm = np.linalg.norm(avg)
if norm > 0:
avg = avg / norm
return avg.astype(np.float32)