RF-DETR / src /infer.py
LosReturn's picture
Upload folder using huggingface_hub
67c84d4 verified
Raw
History Blame Contribute Delete
8.43 kB
from __future__ import annotations
import argparse
import colorsys
from pathlib import Path
import numpy as np
from PIL import Image, ImageDraw, ImageFont
try:
import axengine as ort
BACKEND = "axengine"
print("Running on AXera NPU (axengine)...")
except ImportError:
import onnxruntime as ort
BACKEND = "onnxruntime"
print("Running on CPU/GPU (onnxruntime)...")
MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
CLASSES = [
"person", "bicycle", "car", "motorcycle", "airplane", "bus", "train", "truck", "boat", "traffic_light",
"fire_hydrant", "stop_sign", "parking_meter", "bench", "bird", "cat", "dog", "horse", "sheep", "cow",
"elephant", "bear", "zebra", "giraffe", "backpack", "umbrella", "handbag", "tie", "suitcase", "frisbee",
"skis", "snowboard", "sports_ball", "kite", "baseball_bat", "baseball_glove", "skateboard", "surfboard",
"tennis_racket", "bottle", "wine_glass", "cup", "fork", "knife", "spoon", "bowl", "banana", "apple",
"sandwich", "orange", "broccoli", "carrot", "hot_dog", "pizza", "donut", "cake", "chair", "couch",
"potted_plant", "bed", "dining_table", "toilet", "tv", "laptop", "mouse", "remote", "keyboard",
"cell_phone", "microwave", "oven", "toaster", "sink", "refrigerator", "book", "clock", "vase",
"scissors", "teddy_bear", "hair_drier", "toothbrush",
]
COCO_IDS = [
1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 13, 14, 15, 16, 17, 18, 19, 20, 21,
22, 23, 24, 25, 27, 28, 31, 32, 33, 34, 35, 36, 37, 38, 39, 40, 41, 42, 43, 44,
46, 47, 48, 49, 50, 51, 52, 53, 54, 55, 56, 57, 58, 59, 60, 61, 62, 63, 64, 65,
67, 70, 72, 73, 74, 75, 76, 77, 78, 79, 80, 81, 82, 84, 85, 86, 87, 88, 89, 90,
]
CLASS_NAME_BY_ID = {cid: name for cid, name in zip(COCO_IDS, CLASSES)}
def sigmoid(x: np.ndarray) -> np.ndarray:
return 1.0 / (1.0 + np.exp(-np.clip(x, -88.0, 88.0)))
def get_numpy_dtype(input_meta: object) -> np.dtype:
if hasattr(input_meta, "dtype"):
return np.dtype(input_meta.dtype)
ort_type = getattr(input_meta, "type", "")
mapping = {
"tensor(float)": np.float32,
"tensor(float16)": np.float16,
"tensor(uint8)": np.uint8,
"tensor(int8)": np.int8,
"tensor(int32)": np.int32,
"tensor(int64)": np.int64,
}
if ort_type not in mapping:
raise ValueError(f"Unsupported input type: {ort_type}")
return np.dtype(mapping[ort_type])
def preprocess(
image_path: str,
input_h: int,
input_w: int,
layout: str,
dtype: np.dtype,
letterbox: bool,
) -> tuple[np.ndarray, Image.Image, dict[str, float]]:
raw_image = Image.open(image_path).convert("RGB")
orig_w, orig_h = raw_image.size
if letterbox:
scale_x = min(input_w / orig_w, input_h / orig_h)
scale_y = scale_x
resized_w = max(1, int(round(orig_w * scale_x)))
resized_h = max(1, int(round(orig_h * scale_y)))
pad_x = (input_w - resized_w) // 2
pad_y = (input_h - resized_h) // 2
resized = raw_image.resize((resized_w, resized_h), Image.Resampling.BILINEAR)
canvas = Image.new("RGB", (input_w, input_h), (0, 0, 0))
canvas.paste(resized, (pad_x, pad_y))
image = np.array(canvas)
else:
resized = raw_image.resize((input_w, input_h), Image.Resampling.BILINEAR)
image = np.array(resized)
pad_x = 0.0
pad_y = 0.0
scale_x = input_w / orig_w
scale_y = input_h / orig_h
if BACKEND == "axengine":
if layout == "NHWC":
tensor = image[None, ...].astype(dtype, copy=False)
else:
tensor = image.transpose(2, 0, 1)[None, ...].astype(dtype, copy=False)
else:
if layout == "NHWC":
if dtype != np.uint8:
raise ValueError(f"NHWC input only supports uint8 in this simple script, got {dtype}")
tensor = image[None, ...].astype(np.uint8)
else:
tensor = image.astype(np.float32) / 255.0
tensor = (tensor - MEAN) / STD
tensor = tensor.transpose(2, 0, 1)[None, ...].astype(dtype)
meta = {
"orig_w": float(orig_w),
"orig_h": float(orig_h),
"scale_x": float(scale_x),
"scale_y": float(scale_y),
"pad_x": float(pad_x),
"pad_y": float(pad_y),
"input_w": float(input_w),
"input_h": float(input_h),
}
return tensor, raw_image, meta
def decode(outputs: list[np.ndarray], meta: dict[str, float], thresh: float) -> list[tuple[np.ndarray, float, int, int]]:
dets = outputs[0][0]
labels = outputs[1][0]
if dets.shape[-1] != 4:
dets, labels = labels, dets
boxes = dets
logits = labels[:, :-1]
probs = sigmoid(logits)
input_w = meta["input_w"]
input_h = meta["input_h"]
orig_w = meta["orig_w"]
orig_h = meta["orig_h"]
scale_x = meta["scale_x"]
scale_y = meta["scale_y"]
pad_x = meta["pad_x"]
pad_y = meta["pad_y"]
flat = probs.reshape(-1)
topk = min(300, flat.size)
top_idx = np.argpartition(-flat, topk - 1)[:topk]
top_idx = top_idx[np.argsort(-flat[top_idx])]
num_classes = probs.shape[1]
results = []
for rank, idx in enumerate(top_idx.tolist()):
query_id = idx // num_classes
label_id = idx % num_classes
score = float(flat[idx])
if score < thresh:
continue
cx, cy, bw, bh = boxes[query_id]
x1 = ((cx - bw / 2.0) * input_w - pad_x) / scale_x
y1 = ((cy - bh / 2.0) * input_h - pad_y) / scale_y
x2 = ((cx + bw / 2.0) * input_w - pad_x) / scale_x
y2 = ((cy + bh / 2.0) * input_h - pad_y) / scale_y
x1 = max(0.0, min(orig_w, x1))
y1 = max(0.0, min(orig_h, y1))
x2 = max(0.0, min(orig_w, x2))
y2 = max(0.0, min(orig_h, y2))
results.append((np.array([x1, y1, x2, y2]), score, label_id, query_id))
return results
def color_for_label(label_id: int) -> tuple[int, int, int]:
hue = (label_id * 0.61803398875) % 1.0
r, g, b = colorsys.hsv_to_rgb(hue, 0.75, 1.0)
return int(r * 255), int(g * 255), int(b * 255)
def draw(raw_img: Image.Image, results: list[tuple[np.ndarray, float, int, int]], output_path: str) -> None:
draw_obj = ImageDraw.Draw(raw_img)
try:
font = ImageFont.truetype("DejaVuSans.ttf", 18)
except OSError:
font = ImageFont.load_default()
for box, score, label_id, query_id in results:
x1, y1, x2, y2 = box.tolist()
name = CLASS_NAME_BY_ID.get(label_id, f"obj_{label_id}")
text = f"{name} {score:.2f}"
color = color_for_label(label_id)
draw_obj.rectangle([x1, y1, x2, y2], outline=color, width=3)
draw_obj.rectangle([x1, max(0, y1 - 22), x1 + 140, y1], fill=color)
draw_obj.text((x1 + 2, max(0, y1 - 20)), text, fill="black", font=font)
print(f"query={query_id:3d} class={name:<15} score={score:.4f} box=({x1:.1f}, {y1:.1f}, {x2:.1f}, {y2:.1f})")
Path(output_path).parent.mkdir(parents=True, exist_ok=True)
raw_img.save(output_path)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--model", type=str, required=True)
parser.add_argument("--img", type=str, required=True)
parser.add_argument("--output", type=str, default="result.jpg")
parser.add_argument("--thresh", type=float, default=0.3)
parser.add_argument("--letterbox", action="store_true", help="use letterbox resize instead of direct resize")
args = parser.parse_args()
session = ort.InferenceSession(args.model)
input_meta = session.get_inputs()[0]
shape = [int(x) for x in input_meta.shape]
dtype = get_numpy_dtype(input_meta)
if shape[1] in (1, 3, 4):
layout = "NCHW"
input_h, input_w = shape[2], shape[3]
else:
layout = "NHWC"
input_h, input_w = shape[1], shape[2]
print(f"input_name={input_meta.name} shape={shape} dtype={dtype} layout={layout}")
img_tensor, raw_img, meta = preprocess(args.img, input_h, input_w, layout, dtype, args.letterbox)
outputs = session.run(None, {input_meta.name: img_tensor})
results = decode(outputs, meta, args.thresh)
print(f"Detected {len(results)} objects.")
draw(raw_img, results, args.output)
print(f"Result saved to {args.output}")
if __name__ == "__main__":
main()