spotpredator / app.py
JZVince's picture
revert to original app.py
6d73f1c verified
Raw History Blame Contribute Delete
5.23 kB
"""
app.py — Hugging Face Space to test the SpotPredator PicoDet model in the browser.
Pulls the .tflite straight from the model repo, so testers install nothing — they
drag an image in and see boxes.
Model: JZVince/spotpredator-picodet (PicoDet, Apache-2.0)
PicoDet is exported WITHOUT in-graph NMS, so the TFLite model has TWO outputs:
boxes: (1, N, 4) already-decoded x1,y1,x2,y2 in input-pixel space (0..640)
scores: (1, num_classes, N) per-class confidence
Preprocess is /255 then ImageNet mean/std (RGB) — matches the training pipeline.
"""
import os
import numpy as np
from PIL import Image, ImageDraw
import gradio as gr
from huggingface_hub import hf_hub_download, list_repo_files
MODEL_REPO = "JZVince/spotpredator-picodet"
CLASS_NAMES = ["coyote", "fox", "raptor"] # must match training order / labels.txt
# PicoDet (PaddleDetection) NormalizeImage: /255 then ImageNet mean/std, RGB.
IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
# Optional: set HF_TOKEN as a Space secret for higher rate limits / faster pulls.
HF_TOKEN = os.environ.get("HF_TOKEN")
# ---- fetch the .tflite from the model repo (auto-detect the filename) --------
tflite_name = next(f for f in list_repo_files(MODEL_REPO, token=HF_TOKEN)
if f.endswith(".tflite"))
MODEL_PATH = hf_hub_download(MODEL_REPO, tflite_name, token=HF_TOKEN)
# LiteRT is the successor to tflite-runtime; fall back through older runtimes.
try:
from ai_edge_litert.interpreter import Interpreter
except ImportError:
try:
from tflite_runtime.interpreter import Interpreter
except ImportError:
from tensorflow.lite import Interpreter
interp = Interpreter(model_path=MODEL_PATH)
interp.allocate_tensors()
INP = interp.get_input_details()[0]
OUTS = interp.get_output_details()
IN_W = int(INP["shape"][2])
IN_H = int(INP["shape"][1])
NC = len(CLASS_NAMES)
# PicoDet has two outputs; identify boxes (last dim == 4) vs scores by shape, so
# output order doesn't matter.
BOX_I = 0 if OUTS[0]["shape"][-1] == 4 else 1
SCORE_I = 1 - BOX_I
# ---- NMS (class-aware, applied by caller) -----------------------------------
def nms(boxes, scores, iou_thr=0.45):
x1, y1, x2, y2 = boxes.T
areas = (x2 - x1) * (y2 - y1)
order = scores.argsort()[::-1]
keep = []
while order.size:
i = order[0]
keep.append(i)
xx1 = np.maximum(x1[i], x1[order[1:]])
yy1 = np.maximum(y1[i], y1[order[1:]])
xx2 = np.minimum(x2[i], x2[order[1:]])
yy2 = np.minimum(y2[i], y2[order[1:]])
w = np.maximum(0, xx2 - xx1)
h = np.maximum(0, yy2 - yy1)
inter = w * h
iou = inter / (areas[i] + areas[order[1:]] - inter + 1e-9)
order = order[1:][iou <= iou_thr]
return keep
def detect(image, conf_thr=0.5, iou_thr=0.45):
if image is None:
return None, "Upload an image."
img = image.convert("RGB")
# Preprocess: plain resize to model input, /255, then ImageNet mean/std.
resized = img.resize((IN_W, IN_H), Image.BILINEAR)
x = np.asarray(resized, np.float32) / 255.0
x = ((x - IMAGENET_MEAN) / IMAGENET_STD).astype(np.float32)[None]
interp.set_tensor(INP["index"], x)
interp.invoke()
boxes = interp.get_tensor(OUTS[BOX_I]["index"])[0] # (N, 4) x1,y1,x2,y2 in 0..IN_W
scores = interp.get_tensor(OUTS[SCORE_I]["index"])[0].T # (N, num_classes)
cid = scores.argmax(1)
conf = scores.max(1)
m = conf >= conf_thr
boxes, conf, cid = boxes[m], conf[m], cid[m]
# Scale boxes from model-input space (IN_W x IN_H) back to the original image.
sx = img.width / IN_W
sy = img.height / IN_H
draw = ImageDraw.Draw(img)
colors = [(220, 70, 60), (240, 160, 40), (60, 130, 200)]
lines = []
if len(boxes):
xyxy = boxes.astype(np.float32).copy()
xyxy[:, [0, 2]] *= sx
xyxy[:, [1, 3]] *= sy
xyxy[:, [0, 2]] = xyxy[:, [0, 2]].clip(0, img.width)
xyxy[:, [1, 3]] = xyxy[:, [1, 3]].clip(0, img.height)
for c in np.unique(cid):
idx = np.where(cid == c)[0]
for k in nms(xyxy[idx], conf[idx], iou_thr):
b = xyxy[idx][k]
col = colors[int(c) % 3]
draw.rectangle(list(b), outline=col, width=3)
draw.text((b[0] + 3, max(0, b[1] - 12)),
f"{CLASS_NAMES[int(c)]} {conf[idx][k]:.2f}", fill=col)
lines.append(f"{CLASS_NAMES[int(c)]}: {conf[idx][k]:.2f}")
return img, ("\n".join(lines) if lines else "No predators detected.")
demo = gr.Interface(
fn=detect,
inputs=[gr.Image(type="pil", label="Field image"),
gr.Slider(0.05, 0.9, 0.5, label="Confidence threshold"),
gr.Slider(0.1, 0.9, 0.45, label="NMS IoU")],
outputs=[gr.Image(type="pil", label="Detections"),
gr.Textbox(label="Results")],
title="SpotPredator — PicoDet",
description="Off-grid farm predator detector (coyote / fox / raptor). "
"Runs the PicoDet TFLite model directly (Apache-2.0). Drag in a field image.",
)
if __name__ == "__main__":
demo.launch()