WeDetect 推理测评:精度、速度与显存占用的真实表现

目录

torch推理demo源代码:

精度和速度测评:

服务器封装:

客户端调用:


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))
相关推荐
这张生成的图像能检测吗5 个月前
(论文速读)用于免训练开放词汇表属性检测的组合缓存
目标检测·开放词汇检测
昵称是6硬币10 个月前
SAM3论文精读(逐段解析)
图像分割·sam·实例分割·视觉大模型·sam3·开放词汇检测