scrfd_640_batched / inference_demo.py
ceyxprime's picture
Upload inference_demo.py with huggingface_hub
1f1eb3d verified
Raw
History Blame
2.19 kB
#!/usr/bin/env python3
"""
Minimal demo: batched inference with the fixed det_500m model.
Stacks N input images into a single batch and runs one forward pass.
Prints output shapes and a per-frame summary (count of cls anchors above a
score threshold per scale). Intended as a sanity check, not a full detector —
plug your own NMS / anchor decoding for actual face boxes.
"""
import argparse
import cv2
import numpy as np
import onnxruntime as ort
DET_SIZE = (640, 640)
SCORE_THRESHOLD = 0.5
parser = argparse.ArgumentParser()
parser.add_argument('--model', required=True, help='Path to det_500m_fixed.onnx')
parser.add_argument('--images', nargs='+', required=True,
help='One or more image paths (any number — they form the batch)')
args = parser.parse_args()
def preprocess(path):
img = cv2.imread(path)
if img is None:
raise FileNotFoundError(path)
blob = cv2.dnn.blobFromImage(
cv2.resize(img, DET_SIZE),
1.0 / 128.0, DET_SIZE,
(127.5, 127.5, 127.5), swapRB=True,
)
return blob[0]
# Build batch
batch = np.stack([preprocess(p) for p in args.images], axis=0)
print(f'Input batch: {batch.shape} ({len(args.images)} image(s))')
# One forward pass
sess = ort.InferenceSession(args.model, providers=['CPUExecutionProvider'])
inp_name = sess.get_inputs()[0].name
out_names = [o.name for o in sess.get_outputs()]
outputs = sess.run(None, {inp_name: batch})
# Output shapes
print('\nOutputs:')
for n, o in zip(out_names, outputs):
print(f' {n}: {list(o.shape)}')
# Per-frame summary using the 3 cls heads (post-Sigmoid, indices 0/1/2)
strides = [8, 16, 32]
cls_outputs = outputs[:3]
print(f'\nPer-frame anchor counts above score {SCORE_THRESHOLD}:')
print(f' {"image":<40} {"stride 8":>10} {"stride 16":>10} {"stride 32":>10} {"total":>8}')
for n in range(batch.shape[0]):
counts = [int((cls[n] > SCORE_THRESHOLD).sum()) for cls in cls_outputs]
name = args.images[n]
if len(name) > 38:
name = '...' + name[-35:]
print(f' {name:<40} {counts[0]:>10} {counts[1]:>10} {counts[2]:>10} {sum(counts):>8}')
print('\nIf all images are the same, the per-frame counts must match exactly.')