| |
| """ |
| 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] |
|
|
|
|
| |
| batch = np.stack([preprocess(p) for p in args.images], axis=0) |
| print(f'Input batch: {batch.shape} ({len(args.images)} image(s))') |
|
|
| |
| 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}) |
|
|
| |
| print('\nOutputs:') |
| for n, o in zip(out_names, outputs): |
| print(f' {n}: {list(o.shape)}') |
|
|
| |
| 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.') |
|
|