目录
torch推理demo源代码:
python
import argparse
from typing import List
import torch
import torch.nn.functional as F
import torchvision
from PIL import Image, ImageDraw, ImageFont
# 从你上传的脚本里导入所有模型定义
from test_coco_pytorch import (
XLMRobertaLanguageBackbone,
SimpleYOLOWorldDetector,
load_vision_checkpoint,
)
def build_prompt_embeddings(language_encoder, prompts: List[str], device):
"""把提示词列表编码成 L2 归一化的文本 embedding。"""
with torch.no_grad():
emb = language_encoder(prompts)
emb = F.normalize(emb, dim=-1).to(device)
# 单图推理时 batch=1,保持 (1, K, C) 形状
if emb.dim() == 2:
emb = emb.unsqueeze(0)
return emb
def draw_results(image: Image.Image, result: dict, prompts: List[str]):
"""在 PIL 图片上画框和标签。"""
draw = ImageDraw.Draw(image)
try:
font = ImageFont.truetype("DejaVuSans.ttf", 18)
except Exception:
font = ImageFont.load_default()
boxes = result["bboxes"].cpu()
scores = result["scores"].cpu()
labels = result["labels"].cpu()
for box, score, label in zip(boxes, scores, labels):
x1, y1, x2, y2 = box.tolist()
name = prompts[label.item()]
s = float(score.max().item())
draw.rectangle([x1, y1, x2, y2], outline="red", width=3)
text = f"{name} {s:.2f}"
# 文本背景
bbox = draw.textbbox((x1, max(0, y1 - 20)), text, font=font)
draw.rectangle(bbox, fill="red")
draw.text((x1, max(0, y1 - 20)), text, fill="white", font=font)
return image
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="WeDetect 单图推理")
parser.add_argument("--variant", choices=["tiny", "base", "large"], default="base")
parser.add_argument("--language-model", default="xlm-roberta-base",help="XLM-RoBERTa 模型名或本地路径")
parser.add_argument("--checkpoint", default='assets/wedetect_base.pth',help="wedetect_base.pth 路径")
parser.add_argument("--image", default=r"C:\Users\ChanJing-01\Pictures\890.jpg", help="输入图片路径")
parser.add_argument("--prompts", nargs="+", help="自定义提示词,例如: --prompts 人 汽车 狗")
parser.add_argument("--device", default="cuda")
parser.add_argument("--score-thr", type=float, default=0.01)
parser.add_argument("--nms-iou", type=float, default=0.7)
parser.add_argument("--output", default="output.jpg")
args = parser.parse_args()
args.prompts = ["人"]
device = torch.device(args.device)
# 1) 语言塔:提示词 -> embedding
language_encoder = XLMRobertaLanguageBackbone(
args.language_model, args.checkpoint).to(device).eval()
text_embeddings = build_prompt_embeddings(
language_encoder, args.prompts, device)
if text_embeddings.dim() == 3:
text_embeddings = text_embeddings.squeeze(0)
# 2) 视觉塔 + 检测头
model = SimpleYOLOWorldDetector(
args.variant, score_thr=args.score_thr, nms_iou=args.nms_iou)
load_vision_checkpoint(model, args.checkpoint)
model = model.to(device).eval()
# 3) 单图推理
with torch.no_grad():
results = model([args.image], text_embeddings)
result = results[0]
print(f"检测到 {len(result['bboxes'])} 个目标")
for box, score, label in zip(result["bboxes"], result["scores"],
result["labels"]):
print(f" {args.prompts[label.item()]:<12} "
f"score={float(score.max()):.3f} "
f"box={[int(v) for v in box.tolist()]}")
# 4) 可视化保存
image = Image.open(args.image).convert("RGB")
image = draw_results(image, result, args.prompts)
image.save(args.output)
print(f"结果已保存到 {args.output}")
精度和速度测评:
4060ti上 推理速度50s左右,
召回率比yoloe好
人score=0.055 score=0.055 box=158, 112, 1049, 1775
服务器封装:
python
# api_server.py
import base64
import io
import os
from typing import List
import torch
import torch.nn.functional as F
import uvicorn
from fastapi import FastAPI, File, Form, UploadFile
from fastapi.responses import JSONResponse
from PIL import Image, ImageDraw, ImageFont
from test_coco_pytorch import (
XLMRobertaLanguageBackbone,
SimpleYOLOWorldDetector,
load_vision_checkpoint,
)
# --------------------------------------------------------------------------- #
# 全局模型(启动时加载一次) #
# --------------------------------------------------------------------------- #
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
VARIANT = "base"
LANGUAGE_MODEL = "xlm-roberta-base"
CHECKPOINT = "assets/wedetect_base.pth"
SCORE_THR = 0.01
NMS_IOU = 0.7
OUTPUT_DIR = "outputs"
os.makedirs(OUTPUT_DIR, exist_ok=True)
app = FastAPI(title="WeDetect 开放词汇检测 API")
language_encoder = None
model = None
@app.on_event("startup")
def load_models():
"""启动时加载语言塔和视觉塔,避免每次请求都重新加载。"""
global language_encoder, model
print(f"[startup] device={DEVICE}")
language_encoder = XLMRobertaLanguageBackbone(
LANGUAGE_MODEL, CHECKPOINT).to(DEVICE).eval()
model = SimpleYOLOWorldDetector(
VARIANT, score_thr=SCORE_THR, nms_iou=NMS_IOU)
load_vision_checkpoint(model, CHECKPOINT)
model = model.to(DEVICE).eval()
print("[startup] models loaded")
# --------------------------------------------------------------------------- #
# 工具函数 #
# --------------------------------------------------------------------------- #
def build_prompt_embeddings(prompts: List[str]):
with torch.no_grad():
emb = language_encoder(prompts)
emb = F.normalize(emb, dim=-1).to(DEVICE)
if emb.dim() == 3:
emb = emb.squeeze(0)
return emb
def draw_results(image: Image.Image, result: dict, prompts: List[str]) -> Image.Image:
draw = ImageDraw.Draw(image)
try:
font = ImageFont.truetype("DejaVuSans.ttf", 18)
except Exception:
font = ImageFont.load_default()
boxes = result["bboxes"].cpu()
scores = result["scores"].cpu()
labels = result["labels"].cpu()
for box, score, label in zip(boxes, scores, labels):
x1, y1, x2, y2 = box.tolist()
name = prompts[label.item()]
s = float(score.max().item())
draw.rectangle([x1, y1, x2, y2], outline="red", width=3)
text = f"{name} {s:.2f}"
bbox = draw.textbbox((x1, max(0, y1 - 20)), text, font=font)
draw.rectangle(bbox, fill="red")
draw.text((x1, max(0, y1 - 20)), text, fill="white", font=font)
return image
def image_to_base64(image: Image.Image) -> str:
buf = io.BytesIO()
image.save(buf, format="JPEG", quality=90)
return base64.b64encode(buf.getvalue()).decode("utf-8")
# --------------------------------------------------------------------------- #
# 接口 #
# --------------------------------------------------------------------------- #
@app.get("/health")
def health():
return {"status": "ok", "device": str(DEVICE)}
@app.post("/detect")
async def detect(
file: UploadFile = File(..., description="待检测图片"),
prompts: str = Form(..., description="提示词,逗号分隔,如:人,汽车,狗"),
score_thr: float = Form(SCORE_THR),
nms_iou: float = Form(NMS_IOU),
return_image: bool = Form(False, description="是否返回可视化图片的 base64"),
):
# 1) 解析提示词
prompt_list = [p.strip() for p in prompts.split(",") if p.strip()]
if not prompt_list:
return JSONResponse(status_code=400, content={"error": "prompts 不能为空"})
# 2) 读取图片
try:
img_bytes = await file.read()
image = Image.open(io.BytesIO(img_bytes)).convert("RGB")
except Exception as e:
return JSONResponse(status_code=400, content={"error": f"图片读取失败: {e}"})
# 3) 临时保存图片(模型 forward 接受路径)
tmp_path = os.path.join(OUTPUT_DIR, "_tmp_input.jpg")
image.save(tmp_path)
# 4) 推理
text_embeddings = build_prompt_embeddings(prompt_list)
model.score_thr = score_thr
model.nms_iou = nms_iou
with torch.no_grad():
results = model([tmp_path], text_embeddings)
result = results[0]
# 5) 组装返回
boxes = []
for box, score, label in zip(result["bboxes"], result["scores"],
result["labels"]):
boxes.append({
"label": prompt_list[label.item()],
"score": round(float(score.max().item()), 4),
"box": [int(v) for v in box.tolist()],
})
response = {
"prompts": prompt_list,
"count": len(boxes),
"boxes": boxes,
}
# 6) 可选:可视化图片
if return_image:
vis = draw_results(image.copy(), result, prompt_list)
vis_path = os.path.join(OUTPUT_DIR, "latest_result.jpg")
vis.save(vis_path)
response["image_base64"] = image_to_base64(vis)
return response
if __name__ == "__main__":
uvicorn.run(app, host="0.0.0.0", port=8000)
客户端调用:
dev_client.py
python
import base64
import os
import requests
BASE = "http://127.0.0.1:8000"
IMG = r"C:\Users\ChanJing-01\Pictures\duoshijiao\huizhang.png"
IMG = r"C:\Users\ChanJing-01\Pictures\duoshijiao\021783319821576ab0d945beb4db31a8925a4a25f6e05d9fa8932_0.jpeg"
IMG = r"C:\Users\ChanJing-01\Pictures\duoshijiao\shayu.jpeg"
IMG = r"C:\Users\ChanJing-01\Pictures\jiezhi\jiezhi2.png"
IMG = r"E:\pro_math\math_image\yumaoqiu\imgs\0726_2051_1.jpg"
prompts="娃娃,人,卡通,动物"
prompts="戒指"
save_dir="res"
os.makedirs(save_dir,exist_ok=True)
save_path =save_dir+ "/client_result.jpg"
with open(IMG, "rb") as f:
r = requests.post(
f"{BASE}/detect",
files={"file": f},
data={"prompts": prompts, "score_thr": 0.01, "return_image": True},
)
r.raise_for_status()
resp = r.json()
print("检测到", resp["count"], "个目标")
for b in resp["boxes"]:
print(b)
if "image_base64" in resp:
img_bytes = base64.b64decode(resp["image_base64"])
with open(save_path, "wb") as f:
f.write(img_bytes)
if "image_url" in resp:
img_resp = requests.get(BASE + resp["image_url"])
img_resp.raise_for_status()
with open(save_path, "wb") as f:
f.write(img_resp.content)
print("图片已保存到", os.path.abspath(save_path))