
目录torch推理demo源代码精度和速度测评服务器封装客户端调用torch推理demo源代码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) # 单图推理时 batch1保持 (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], outlinered, width3) text f{name} {s:.2f} # 文本背景 bbox draw.textbbox((x1, max(0, y1 - 20)), text, fontfont) draw.rectangle(bbox, fillred) draw.text((x1, max(0, y1 - 20)), text, fillwhite, fontfont) return image if __name__ __main__: parser argparse.ArgumentParser(descriptionWeDetect 单图推理) parser.add_argument(--variant, choices[tiny, base, large], defaultbase) parser.add_argument(--language-model, defaultxlm-roberta-base,helpXLM-RoBERTa 模型名或本地路径) parser.add_argument(--checkpoint, defaultassets/wedetect_base.pth,helpwedetect_base.pth 路径) parser.add_argument(--image, defaultrC:\Users\ChanJing-01\Pictures\890.jpg, help输入图片路径) parser.add_argument(--prompts, nargs, help自定义提示词例如: --prompts 人 汽车 狗) parser.add_argument(--device, defaultcuda) parser.add_argument(--score-thr, typefloat, default0.01) parser.add_argument(--nms-iou, typefloat, default0.7) parser.add_argument(--output, defaultoutput.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_thrargs.score_thr, nms_iouargs.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} fscore{float(score.max()):.3f} fbox{[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好人score0.055 score0.055 box[158, 112, 1049, 1775]服务器封装# 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_okTrue) app FastAPI(titleWeDetect 开放词汇检测 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_thrSCORE_THR, nms_iouNMS_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], outlinered, width3) text f{name} {s:.2f} bbox draw.textbbox((x1, max(0, y1 - 20)), text, fontfont) draw.rectangle(bbox, fillred) draw.text((x1, max(0, y1 - 20)), text, fillwhite, fontfont) return image def image_to_base64(image: Image.Image) - str: buf io.BytesIO() image.save(buf, formatJPEG, quality90) 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_code400, 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_code400, 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, host0.0.0.0, port8000)客户端调用dev_client.pyimport base64 import os import requests BASE http://127.0.0.1:8000 IMG rC:\Users\ChanJing-01\Pictures\duoshijiao\huizhang.png IMG rC:\Users\ChanJing-01\Pictures\duoshijiao\021783319821576ab0d945beb4db31a8925a4a25f6e05d9fa8932_0.jpeg IMG rC:\Users\ChanJing-01\Pictures\duoshijiao\shayu.jpeg IMG rC:\Users\ChanJing-01\Pictures\jiezhi\jiezhi2.png IMG rE:\pro_math\math_image\yumaoqiu\imgs\0726_2051_1.jpg prompts娃娃,人卡通动物 prompts戒指 save_dirres os.makedirs(save_dir,exist_okTrue) 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))