GeCo2:基于渐进式查询聚合的通用尺度目标计数模型部署实践
·
说明
地址:https://github.com/jerpelhan/GECO2

效果

代码
import tkinter as tk
from tkinter import filedialog, messagebox
from PIL import Image, ImageTk
import torch
from torchvision.ops import nms
from torchvision import transforms as T
from utils.arg_parser import get_argparser
from models.counter_infer import build_model
from utils.data import resize_and_pad
MODEL_PATH = "CNTQG_multitrain_ca44.pth"
DEVICE = torch.device("cuda"if torch.cuda.is_available() else"cpu")
IMG_SIZE = 1024
THRESHOLD = 0.33 # 阈值越小框越多
NMS_IOU = 0.5
class GECOApp:
def __init__(self, root):
self.root = root
self.root.title("GECO2 计数检测")
self.root.geometry("900x700")
self.model = self._load_model()
self.original_image = None
self.display_image = None
self.tk_image = None
self.user_boxes = []
self.start_x = None
self.start_y = None
self.drawing = False
self.scale_disp = 1.0
self.offset_x = 0
self.offset_y = 0
self.build_ui()
def _load_model(self):
print("[初始化] 加载模型...")
args = get_argparser().parse_args([])
args.zero_shot = True
model = build_model(args).to(DEVICE)
state_dict = torch.load(MODEL_PATH, map_location=DEVICE)['model']
state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}
model.load_state_dict(state_dict, strict=False)
model.eval()
model.return_masks = False
print("[初始化] 模型加载完成")
return model
def build_ui(self):
toolbar = tk.Frame(self.root)
toolbar.pack(side=tk.TOP, fill=tk.X, padx=5, pady=5)
self.btn_open = tk.Button(toolbar, text="打开图片", command=self.load_image)
self.btn_open.pack(side=tk.LEFT, padx=5)
self.btn_clear = tk.Button(toolbar, text="清除画框", command=self.clear_boxes)
self.btn_clear.pack(side=tk.LEFT, padx=5)
self.btn_detect = tk.Button(toolbar, text="开始计数", command=self.detect_objects, state=tk.DISABLED)
self.btn_detect.pack(side=tk.LEFT, padx=5)
self.lbl_result = tk.Label(toolbar, text="计数: -", fg="blue", font=("Arial", 14, "bold"))
self.lbl_result.pack(side=tk.RIGHT, padx=20)
self.canvas = tk.Canvas(self.root, bg='gray')
self.canvas.pack(fill=tk.BOTH, expand=True)
self.canvas.bind("<ButtonPress-1>", self.on_mouse_down)
self.canvas.bind("<B1-Motion>", self.on_mouse_move)
self.canvas.bind("<ButtonRelease-1>", self.on_mouse_up)
self.status = tk.Label(self.root, text="就绪", bd=1, relief=tk.SUNKEN, anchor=tk.W)
self.status.pack(side=tk.BOTTOM, fill=tk.X)
def load_image(self):
path = filedialog.askopenfilename(filetypes=[("图片", "*.jpg *.jpeg *.png *.bmp")])
if not path:
return
self.original_image = Image.open(path).convert("RGB")
self.user_boxes = []
print(f"[打开图片] 文件: {path}, 尺寸: {self.original_image.size}")
self._redraw()
self.btn_detect.config(state=tk.NORMAL if self.user_boxes else tk.DISABLED)
def clear_boxes(self):
print("[清除] 移除所有示例框")
self.user_boxes = []
self._redraw()
self.btn_detect.config(state=tk.DISABLED)
def _resize_to_canvas(self, pil_img):
cw = self.canvas.winfo_width()
ch = self.canvas.winfo_height()
if cw < 10 or ch < 10:
cw, ch = 800, 600
iw, ih = pil_img.size
scale = min(cw / iw, ch / ih)
new_w, new_h = int(iw * scale), int(ih * scale)
return pil_img.resize((new_w, new_h), Image.LANCZOS), scale
def _redraw(self, pred_boxes=None):
self.canvas.delete("all")
if self.original_image is None:
return
display_pil, scale = self._resize_to_canvas(self.original_image)
self.scale_disp = scale
self.tk_image = ImageTk.PhotoImage(display_pil)
cw = self.canvas.winfo_width()
ch = self.canvas.winfo_height()
dw, dh = display_pil.size
self.offset_x = (cw - dw) // 2
self.offset_y = (ch - dh) // 2
self.canvas.create_image(self.offset_x, self.offset_y, anchor=tk.NW, image=self.tk_image)
for (x1, y1, x2, y2) in self.user_boxes:
dx1 = x1 * scale + self.offset_x
dy1 = y1 * scale + self.offset_y
dx2 = x2 * scale + self.offset_x
dy2 = y2 * scale + self.offset_y
self.canvas.create_rectangle(dx1, dy1, dx2, dy2, outline="red", width=2)
if pred_boxes:
for (x1, y1, x2, y2) in pred_boxes:
dx1 = x1 * scale + self.offset_x
dy1 = y1 * scale + self.offset_y
dx2 = x2 * scale + self.offset_x
dy2 = y2 * scale + self.offset_y
self.canvas.create_rectangle(dx1, dy1, dx2, dy2, outline="orange", width=2)
# ---------- 鼠标事件 ----------
def on_mouse_down(self, event):
if self.original_image is None:
return
x = (event.x - self.offset_x) / self.scale_disp
y = (event.y - self.offset_y) / self.scale_disp
if 0 <= x <= self.original_image.width and 0 <= y <= self.original_image.height:
self.start_x, self.start_y = x, y
self.drawing = True
print(f"[鼠标按下] 原图坐标: ({x:.1f}, {y:.1f})")
def on_mouse_move(self, event):
if not self.drawing:
return
self._redraw()
x = (event.x - self.offset_x) / self.scale_disp
y = (event.y - self.offset_y) / self.scale_disp
dx1 = self.start_x * self.scale_disp + self.offset_x
dy1 = self.start_y * self.scale_disp + self.offset_y
dx2 = x * self.scale_disp + self.offset_x
dy2 = y * self.scale_disp + self.offset_y
self.canvas.create_rectangle(dx1, dy1, dx2, dy2, outline="blue", dash=(3, 5))
def on_mouse_up(self, event):
if not self.drawing:
return
self.drawing = False
x = (event.x - self.offset_x) / self.scale_disp
y = (event.y - self.offset_y) / self.scale_disp
if abs(x - self.start_x) < 5 or abs(y - self.start_y) < 5:
print(f"[鼠标释放] 矩形过小,忽略")
self._redraw()
return
x1, y1 = min(self.start_x, x), min(self.start_y, y)
x2, y2 = max(self.start_x, x), max(self.start_y, y)
self.user_boxes.append([x1, y1, x2, y2])
print(f"[新增示例框] 原图坐标: ({x1:.1f}, {y1:.1f}, {x2:.1f}, {y2:.1f}), 当前总数: {len(self.user_boxes)}")
self._redraw()
self.btn_detect.config(state=tk.NORMAL)
# ---------- 推理 ----------
def detect_objects(self):
if not self.user_boxes:
messagebox.showwarning("提示", "请先绘制至少一个示例框!")
return
try:
self.status.config(text="推理中...")
self.root.update()
# 1. 图像转 Tensor
img_tensor = T.ToTensor()(self.original_image).to(DEVICE)
boxes_tensor = torch.tensor(self.user_boxes, dtype=torch.float32, device=DEVICE)
print(f"[推理] 原始图像 tensor 形状: {img_tensor.shape}, 示例框: {boxes_tensor.shape}")
# 2. resize_and_pad
img_padded, boxes_scaled, scale = resize_and_pad(img_tensor, boxes_tensor, size=float(IMG_SIZE))
print(f"[推理 resize_and_pad] 缩放因子: {scale:.4f}")
print(f"[推理] padded 图像形状: {img_padded.shape}, 缩放后示例框: {boxes_scaled}")
# 3. 标准化
img_norm = T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])(img_padded)
print(f"[推理] 标准化后图像统计: min={img_norm.min():.3f}, max={img_norm.max():.3f}, mean={img_norm.mean():.3f}")
img_batch = img_norm.unsqueeze(0)
boxes_batch = boxes_scaled.unsqueeze(0)
# 4. 模型前向
with torch.no_grad():
outputs, _, _, _, _ = self.model(img_batch, boxes_batch)
out = outputs[0]
pred_boxes = out['pred_boxes']
if isinstance(pred_boxes, list):
pred_boxes = pred_boxes[-1]
if pred_boxes.dim() == 3:
pred_boxes = pred_boxes.squeeze(0)
box_v = out['box_v']
if isinstance(box_v, list):
box_v = box_v[-1]
if box_v.dim() == 2:
box_v = box_v.squeeze(0)
print(f"[推理输出] 原始预测框数量: {pred_boxes.shape[0]}")
if pred_boxes.shape[0] > 0:
print(f"[推理输出] box_v 范围: {box_v.min():.4f} ~ {box_v.max():.4f}")
# 5. 阈值筛选
if pred_boxes.numel() == 0:
final_boxes = []
count = 0
else:
thr = 1.0 / 0.33 # 与原demo一致的阈值处理
mask = box_v > (box_v.max() / thr)
pred_boxes = pred_boxes[mask]
box_v = box_v[mask]
print(f"[阈值筛选] 后框数: {pred_boxes.shape[0]}")
# 6. NMS
if pred_boxes.shape[0] > 0:
keep = nms(pred_boxes, box_v, NMS_IOU)
pred_boxes = pred_boxes[keep]
box_v = box_v[keep]
print(f"[NMS] 后框数: {len(keep)}")
else:
print("[NMS] 无框,跳过")
# 7. 坐标映射
if pred_boxes.shape[0] > 0:
pred_boxes_pix = pred_boxes * IMG_SIZE
pred_boxes_orig = pred_boxes_pix / scale
final_boxes = pred_boxes_orig.tolist()
count = len(final_boxes)
print(f"[坐标映射] 最终检测到 {count} 个目标")
for i, box in enumerate(final_boxes):
print(f" 框{i}: [{box[0]:.1f}, {box[1]:.1f}, {box[2]:.1f}, {box[3]:.1f}]")
else:
final_boxes = []
count = 0
print("[坐标映射] 无最终框")
# 更新界面
self._redraw(pred_boxes=final_boxes)
self.lbl_result.config(text=f"计数: {count}")
self.status.config(text=f"检测完成,共 {count} 个目标")
except Exception as e:
messagebox.showerror("错误", f"推理失败:\n{str(e)}")
self.status.config(text="推理出错")
raise
if __name__ == "__main__":
root = tk.Tk()
app = GECOApp(root)
root.mainloop()
更多推荐




所有评论(0)