yolo_detector.py 20 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475
  1. # -*- coding: utf-8 -*-
  2. """
  3. YOLO 卡牌检测裁切模块
  4. 默认查询路径(2026-08):card_seg_v2 (yolo26s-seg) → mask 4 点 → warp_perspective → ensure_portrait
  5. - 向量库/图库特征不动;仅优化「待匹配拍摄图」的裁切与几何摆正
  6. - TRT:CARD_SEG_USE_TRT=1 时优先加载同目录同名 .engine(.pt/.onnx 均可推导)
  7. - 回退:CARD_CROP_MODE=bbox_legacy 或 CARD_CROP_MODE 未设且仍指向旧 onnx 时走矩形 bbox 裁切
  8. 环境变量:
  9. CARD_SEG_USE_TRT=1 启用 TensorRT engine(须同目录存在 .engine)
  10. CARD_CROP_MODE=seg_warp|bbox_legacy 覆盖 config.CARD_CROP_MODE
  11. CARD_MERGE_DIST_PX 近距 mask 合并距离阈值(px),<=0 关闭;默认取 config(50)
  12. CARD_MERGE_MIN_IOU 合并可选 IoU 门槛,0=纯距离口径;默认取 config(0)
  13. """
  14. import math
  15. import os
  16. import numpy as np
  17. import cv2
  18. from modules.ultralytics_compat import import_yolo
  19. def _env_truthy(name, default="0"):
  20. return os.environ.get(name, default).lower() in ("1", "true", "yes", "seg_warp")
  21. def _resolve_engine_path(model_path):
  22. """同目录同名 .engine:xxx.pt / xxx.onnx → xxx.engine"""
  23. if model_path.endswith(".onnx"):
  24. return model_path[:-5] + ".engine"
  25. if model_path.endswith(".pt"):
  26. return model_path[:-3] + ".engine"
  27. if model_path.endswith(".engine"):
  28. return model_path
  29. return None
  30. def _pack_crop(crop, box=None, label=None, return_box=False, return_label=False):
  31. """按需打包 detect_and_crop 返回值:crop / (crop,box) / (crop,label) / (crop,box,label)。"""
  32. if return_label and return_box:
  33. return crop, box, label
  34. if return_label:
  35. return crop, label
  36. if return_box:
  37. return crop, box
  38. return crop
  39. def _sort_reading_order(items):
  40. """多卡按阅读顺序:先上后下分行,同行从左到右。
  41. items 每项需含 box=(x1,y1,x2,y2)。"""
  42. if not items:
  43. return items
  44. heights = [max(1.0, it["box"][3] - it["box"][1]) for it in items]
  45. row_tol = 0.5 * float(np.median(heights))
  46. ordered = sorted(
  47. items,
  48. key=lambda it: ((it["box"][1] + it["box"][3]) / 2.0,
  49. (it["box"][0] + it["box"][2]) / 2.0),
  50. )
  51. rows = []
  52. for it in ordered:
  53. cy = (it["box"][1] + it["box"][3]) / 2.0
  54. if rows and abs(cy - rows[-1]["cy"]) <= row_tol:
  55. rows[-1]["items"].append(it)
  56. n = len(rows[-1]["items"])
  57. rows[-1]["cy"] = (rows[-1]["cy"] * (n - 1) + cy) / n
  58. else:
  59. rows.append({"cy": cy, "items": [it]})
  60. out = []
  61. for row in rows:
  62. out.extend(sorted(row["items"],
  63. key=lambda it: (it["box"][0] + it["box"][2]) / 2.0))
  64. return out
  65. def _box_iou(a, b):
  66. """xyxy 框 IoU(近距合并的可选门槛用)。"""
  67. ix1, iy1 = max(a[0], b[0]), max(a[1], b[1])
  68. ix2, iy2 = min(a[2], b[2]), min(a[3], b[3])
  69. iw, ih = max(0.0, ix2 - ix1), max(0.0, iy2 - iy1)
  70. inter = iw * ih
  71. if inter <= 0:
  72. return 0.0
  73. area_a = max(0.0, a[2] - a[0]) * max(0.0, a[3] - a[1])
  74. area_b = max(0.0, b[2] - b[0]) * max(0.0, b[3] - b[1])
  75. union = area_a + area_b - inter
  76. return inter / union if union > 0 else 0.0
  77. def _center_of(mask, box, min_area_frac=0.01, stride=4):
  78. """实例中心:mask 质心(有效面积 >= 1% 图幅,防噪声拉偏),否则 bbox 中心兜底。
  79. 均为原图分辨率坐标 (cx, cy)。质心用 stride 跨步采样:retina mask 是原图分辨率
  80. (4256×2128 全图 nonzero 实测 ~10ms/实例,占 8020 post 段大头),采样后
  81. ±stride px 抖动远小于 50px 合并阈值,合并/排序行为不变。"""
  82. if mask is not None:
  83. sub = mask[::stride, ::stride] > 0.5
  84. ys, xs = np.nonzero(sub)
  85. if xs.size >= min_area_frac * sub.shape[0] * sub.shape[1]:
  86. return (float(xs.mean()) * stride, float(ys.mean()) * stride)
  87. return ((box[0] + box[2]) / 2.0, (box[1] + box[3]) / 2.0)
  88. class CardDetector:
  89. """卡牌区域检测 + 查询侧裁切(bbox 或 seg+warp)"""
  90. def __init__(self, model_path, conf=0.25, imgsz=640, device=None, crop_mode=None):
  91. YOLO = import_yolo()
  92. if not os.path.exists(model_path):
  93. raise FileNotFoundError(f"YOLO 模型不存在: {model_path}")
  94. self.conf = conf
  95. self.imgsz = imgsz
  96. self.model_path = model_path
  97. # crop_mode: seg_warp | bbox_legacy
  98. env_mode = os.environ.get("CARD_CROP_MODE", "").strip().lower()
  99. if crop_mode is not None:
  100. self.crop_mode = crop_mode
  101. elif env_mode in ("seg_warp", "bbox_legacy"):
  102. self.crop_mode = env_mode
  103. else:
  104. try:
  105. import config as _cfg
  106. self.crop_mode = getattr(_cfg, "CARD_CROP_MODE", "seg_warp")
  107. except Exception:
  108. self.crop_mode = "seg_warp"
  109. if self.crop_mode not in ("seg_warp", "bbox_legacy"):
  110. self.crop_mode = "seg_warp"
  111. # 近距 mask 合并阈值:env → config → 默认(detect_and_crop_all 用;<=0 关闭合并)
  112. def _merge_cfg(key, default):
  113. v = os.environ.get(key, "").strip()
  114. if v:
  115. try:
  116. return float(v)
  117. except ValueError:
  118. pass
  119. try:
  120. import config as _cfg
  121. return float(getattr(_cfg, key, default))
  122. except Exception:
  123. return float(default)
  124. self.merge_dist_px = _merge_cfg("CARD_MERGE_DIST_PX", 50.0)
  125. self.merge_min_iou = _merge_cfg("CARD_MERGE_MIN_IOU", 0.0)
  126. self._trt = False
  127. use_trt = os.environ.get("CARD_SEG_USE_TRT", "0").lower() in ("1", "true", "yes")
  128. engine_path = _resolve_engine_path(model_path) if use_trt else None
  129. # 若直接传入 .engine
  130. if model_path.endswith(".engine") and os.path.exists(model_path):
  131. engine_path = model_path
  132. use_trt = True
  133. if use_trt and engine_path and os.path.exists(engine_path):
  134. try:
  135. import tensorrt # noqa: F401
  136. self.model = YOLO(engine_path, task="segment")
  137. self.device = device if device is not None else 0
  138. self._trt = True
  139. print(
  140. f"[CardDetector] TRT engine mode={self.crop_mode} device={self.device} "
  141. f"engine={os.path.basename(engine_path)}",
  142. flush=True,
  143. )
  144. except Exception as e:
  145. print(
  146. f"[CardDetector] TRT 加载失败({type(e).__name__}: {e}), 回退 {model_path}",
  147. flush=True,
  148. )
  149. self.model = YOLO(model_path, task="segment")
  150. self.device = device
  151. else:
  152. self.model = YOLO(model_path, task="segment")
  153. self.device = device
  154. if use_trt and engine_path and not os.path.exists(engine_path):
  155. print(
  156. f"[CardDetector] CARD_SEG_USE_TRT=1 但 engine 不存在: {engine_path}, 用 {model_path}",
  157. flush=True,
  158. )
  159. def _predict(self, img):
  160. is_path = isinstance(img, str)
  161. results = self.model.predict(
  162. img if is_path else img,
  163. conf=self.conf,
  164. imgsz=self.imgsz,
  165. device=self.device,
  166. verbose=False,
  167. retina_masks=True,
  168. )
  169. return results, is_path
  170. @staticmethod
  171. def _top_label(res, idx):
  172. """取第 idx 个实例的类别名(yolo26s-seg 的 label,如 'pokemon'/'nba');取不到返回 None。"""
  173. try:
  174. boxes = res.boxes
  175. if boxes is None or boxes.cls is None or len(boxes.cls) <= idx:
  176. return None
  177. cls_id = int(boxes.cls[idx].item() if hasattr(boxes.cls[idx], "item") else boxes.cls[idx])
  178. return (res.names or {}).get(cls_id)
  179. except Exception:
  180. return None
  181. def detect_and_crop(self, img, return_box=False, return_label=False):
  182. """
  183. 检测并裁切最大卡牌区域。
  184. - crop_mode=bbox_legacy: 仅 xyxy 矩形裁切(旧行为)
  185. - crop_mode=seg_warp: mask→4点→warp(valid_quad 则始终 warp)→ensure_portrait
  186. 返回 RGB;失败时行为与旧版一致(尽量返回原图 RGB,return_box 时 box 可能为 None)
  187. return_label=True 时额外返回最大实例的 yolo label 名称(如 'pokemon'):
  188. (crop, label) 或 (crop, box, label);未检出卡时 label=None
  189. """
  190. if self.crop_mode == "seg_warp":
  191. return self._detect_and_crop_seg_warp(img, return_box=return_box, return_label=return_label)
  192. return self._detect_and_crop_bbox(img, return_box=return_box, return_label=return_label)
  193. def detect_and_crop_all(self, img):
  194. """检测图中全部卡牌实例,按阅读顺序(上→下、同行左→右)裁切。
  195. YOLO 只跑一次。近距合并:两实例中心(mask 质心,回退 bbox 中心)距离
  196. <= merge_dist_px 视为同一张卡的重复分割,仅保留 conf 高者(8020 / 8000
  197. 多卡路径共用本逻辑;阈值 CARD_MERGE_DIST_PX 可 env/config 调,<=0 关闭)。
  198. 返回 list[{"crop": RGB, "box": (x1,y1,x2,y2), "label": str|None,
  199. "conf": float, "center": (cx,cy)}],conf/center 为新增键(加法,
  200. query_all 等既有调用方不受影响)。
  201. 未检出返回 [](与 detect_and_crop 单卡失败回原图的行为刻意不同——多卡路径不能把整图当一张卡)。"""
  202. is_path = isinstance(img, str)
  203. results, _ = self._predict(img)
  204. if not results:
  205. return []
  206. res = results[0]
  207. orig = res.orig_img
  208. if orig is None:
  209. return []
  210. h, w = orig.shape[:2]
  211. orig_rgb = cv2.cvtColor(orig, cv2.COLOR_BGR2RGB) if is_path else orig
  212. boxes = res.boxes
  213. if boxes is None or len(boxes) == 0:
  214. return []
  215. xyxy = boxes.xyxy.cpu().numpy()
  216. confs = None
  217. try:
  218. if boxes.conf is not None and len(boxes.conf) == len(xyxy):
  219. confs = boxes.conf.cpu().numpy()
  220. except Exception:
  221. confs = None
  222. masks_np = None
  223. if res.masks is not None:
  224. masks_np = res.masks.data.cpu().numpy()
  225. # 第一段:逐实例元信息(跨步采样质心 + conf;全分辨率二值化推迟到幸存者,
  226. # 被吸收实例不再白做二值化——post 段 2026-09-03 profile 占单请求 ~11%)
  227. metas = []
  228. for i in range(len(xyxy)):
  229. box = tuple(float(v) for v in xyxy[i])
  230. mask_raw = None
  231. if masks_np is not None and i < len(masks_np):
  232. mask_raw = masks_np[i]
  233. if mask_raw.shape[:2] != (h, w):
  234. mask_raw = cv2.resize(mask_raw, (w, h), interpolation=cv2.INTER_NEAREST)
  235. metas.append({
  236. "box": box,
  237. "label": self._top_label(res, i),
  238. "mask_raw": mask_raw,
  239. "conf": float(confs[i]) if confs is not None else 0.0,
  240. "center": _center_of(mask_raw, box),
  241. })
  242. # 第二段:近距合并(同一张卡被拆成多个 mask / 重复检测 → 保 conf 高者)
  243. keep, _dropped = self._merge_indices(metas)
  244. # 第三段:只对幸存者二值化全分辨率 mask + 裁卡(被吸收者连二值化都不做)
  245. items = []
  246. for m in keep:
  247. mask = None
  248. if m["mask_raw"] is not None:
  249. mask = (m["mask_raw"] > 0.5).astype(np.uint8) * 255
  250. crop = self._crop_one_instance(orig_rgb, m["box"], mask)
  251. items.append({
  252. "crop": crop,
  253. "box": tuple(int(round(v)) for v in m["box"]),
  254. "label": m["label"],
  255. "conf": m["conf"],
  256. "center": m["center"],
  257. })
  258. return _sort_reading_order(items)
  259. def _merge_indices(self, metas):
  260. """近距合并:conf 降序贪心——每个实例与已保留实例算中心欧氏距,
  261. 距离 <= merge_dist_px(且 IoU >= merge_min_iou,若门槛>0 启用)则被吸收。
  262. 返回 (保留 metas 列表(按原顺序), 被吸收索引列表)。"""
  263. thr = self.merge_dist_px
  264. min_iou = self.merge_min_iou
  265. if thr <= 0 or len(metas) <= 1:
  266. return metas, []
  267. order = sorted(range(len(metas)), key=lambda i: -metas[i]["conf"])
  268. kept, dropped = [], []
  269. for i in order:
  270. mi = metas[i]
  271. dup = False
  272. for j in kept:
  273. mj = metas[j]
  274. d = math.hypot(mi["center"][0] - mj["center"][0],
  275. mi["center"][1] - mj["center"][1])
  276. if d > thr:
  277. continue
  278. if min_iou > 0 and _box_iou(mi["box"], mj["box"]) < min_iou:
  279. continue
  280. dup = True
  281. break
  282. if dup:
  283. dropped.append(i)
  284. else:
  285. kept.append(i)
  286. if dropped:
  287. print(f"[CardDetector] 近距合并(thr={thr:g}px"
  288. + (f" iou>={min_iou:g}" if min_iou > 0 else "")
  289. + f"): 保留 {sorted(kept)} 吸收 {sorted(dropped)} "
  290. + str([(metas[i]["center"], round(metas[i]["conf"], 3)) for i in dropped]),
  291. flush=True)
  292. return [metas[i] for i in sorted(kept)], dropped
  293. def _crop_one_instance(self, orig_rgb, box, mask):
  294. """单实例 → 裁卡 RGB(seg_warp 优先,失败回退 bbox)。"""
  295. from modules.rectifier import (
  296. mask_to_corners, order_corners, valid_quad, warp_perspective,
  297. crop_bbox, ensure_portrait,
  298. )
  299. out = None
  300. if self.crop_mode == "seg_warp" and mask is not None:
  301. corners = mask_to_corners(mask)
  302. if corners is not None:
  303. ordered = order_corners(corners)
  304. if valid_quad(ordered):
  305. warped = warp_perspective(orig_rgb, corners)
  306. if warped is not None and warped.size > 0:
  307. out = warped
  308. if out is None and box is not None:
  309. out = crop_bbox(orig_rgb, box)
  310. if out is None:
  311. out = orig_rgb
  312. out = ensure_portrait(out)
  313. if out.dtype != np.uint8:
  314. out = np.clip(out, 0, 255).astype(np.uint8)
  315. return out
  316. def _detect_and_crop_bbox(self, img, return_box=False, return_label=False):
  317. is_path = isinstance(img, str)
  318. results, _ = self._predict(img)
  319. if not results:
  320. return _pack_crop(None, return_box=return_box, return_label=return_label)
  321. res = results[0]
  322. orig = res.orig_img # 路径输入多为 BGR;数组输入约定与 ultralytics 一致
  323. if orig is None:
  324. return _pack_crop(None, return_box=return_box, return_label=return_label)
  325. boxes = res.boxes
  326. if boxes is None or len(boxes) == 0:
  327. rgb = cv2.cvtColor(orig, cv2.COLOR_BGR2RGB) if is_path else orig
  328. return _pack_crop(rgb, return_box=return_box, return_label=return_label)
  329. xyxy = boxes.xyxy.cpu().numpy()
  330. areas = (xyxy[:, 2] - xyxy[:, 0]) * (xyxy[:, 3] - xyxy[:, 1])
  331. max_idx = int(np.argmax(areas))
  332. label = self._top_label(res, max_idx)
  333. x1, y1, x2, y2 = xyxy[max_idx]
  334. h, w = orig.shape[:2]
  335. x1, y1 = max(0, int(x1)), max(0, int(y1))
  336. x2, y2 = min(w, int(x2)), min(h, int(y2))
  337. if x2 <= x1 or y2 <= y1:
  338. rgb = cv2.cvtColor(orig, cv2.COLOR_BGR2RGB) if is_path else orig
  339. return _pack_crop(rgb, return_box=return_box, return_label=return_label)
  340. crop = orig[y1:y2, x1:x2]
  341. crop_rgb = cv2.cvtColor(crop, cv2.COLOR_BGR2RGB) if is_path else crop
  342. return _pack_crop(crop_rgb, box=(x1, y1, x2, y2), label=label,
  343. return_box=return_box, return_label=return_label)
  344. def _detect_and_crop_seg_warp(self, img, return_box=False, return_label=False):
  345. """seg → 最大实例 mask → 4 点 → warp(始终尝试)→ ensure_portrait → RGB"""
  346. from modules.rectifier import (
  347. mask_to_corners, order_corners, valid_quad, warp_perspective,
  348. crop_bbox, ensure_portrait,
  349. )
  350. is_path = isinstance(img, str)
  351. if return_label:
  352. box, mask, orig_rgb, label = self.detect_box_mask(img, return_label=True)
  353. else:
  354. box, mask, orig_rgb = self.detect_box_mask(img)
  355. label = None
  356. if orig_rgb is None:
  357. return _pack_crop(None, return_box=return_box, return_label=return_label)
  358. # 在 RGB 上 warp,与下游 DINOv2 / letterbox 一致
  359. out = None
  360. if mask is not None:
  361. corners = mask_to_corners(mask)
  362. if corners is not None:
  363. ordered = order_corners(corners)
  364. # 查询侧:valid_quad 即 warp,不再用 has_perspective 跳过
  365. if valid_quad(ordered):
  366. warped = warp_perspective(orig_rgb, corners)
  367. if warped is not None and warped.size > 0:
  368. out = warped
  369. if out is None and box is not None:
  370. out = crop_bbox(orig_rgb, box)
  371. if out is None:
  372. out = orig_rgb
  373. out = ensure_portrait(out)
  374. # 保证 uint8 RGB
  375. if out.dtype != np.uint8:
  376. out = np.clip(out, 0, 255).astype(np.uint8)
  377. box_i = None
  378. if box is not None:
  379. box_i = tuple(int(round(v)) for v in box)
  380. return _pack_crop(out, box=box_i, label=label,
  381. return_box=return_box, return_label=return_label)
  382. def detect_box_mask(self, img, return_label=False):
  383. """
  384. 检测最大卡牌,返回 (box, mask, orig_rgb) 供透视矫正;return_label=True 时
  385. 返回 (box, mask, orig_rgb, label)(label 为 yolo 类别名,未检出时 None)。
  386. 颜色约定:路径输入 → RGB;数组输入原样返回(调用方保证语义)。
  387. """
  388. is_path = isinstance(img, str)
  389. results, _ = self._predict(img)
  390. if not results:
  391. return (None, None, None, None) if return_label else (None, None, None)
  392. res = results[0]
  393. orig = res.orig_img
  394. if orig is None:
  395. return (None, None, None, None) if return_label else (None, None, None)
  396. h, w = orig.shape[:2]
  397. # ultralytics 路径读图为 BGR;numpy 输入时 orig 通常与输入同通道序
  398. if is_path:
  399. orig_rgb = cv2.cvtColor(orig, cv2.COLOR_BGR2RGB)
  400. else:
  401. orig_rgb = orig
  402. boxes = res.boxes
  403. if boxes is None or len(boxes) == 0:
  404. return (None, None, orig_rgb, None) if return_label else (None, None, orig_rgb)
  405. xyxy = boxes.xyxy.cpu().numpy()
  406. areas = (xyxy[:, 2] - xyxy[:, 0]) * (xyxy[:, 3] - xyxy[:, 1])
  407. idx = int(np.argmax(areas))
  408. box = tuple(float(v) for v in xyxy[idx])
  409. label = self._top_label(res, idx)
  410. mask = None
  411. if res.masks is not None:
  412. m = res.masks.data.cpu().numpy()[idx]
  413. m = (m > 0.5).astype(np.uint8) * 255
  414. if m.shape[:2] != (h, w):
  415. m = cv2.resize(m, (w, h), interpolation=cv2.INTER_NEAREST)
  416. mask = m
  417. if return_label:
  418. return box, mask, orig_rgb, label
  419. return box, mask, orig_rgb
  420. def crop_batch(self, img_list):
  421. out = []
  422. for img in img_list:
  423. try:
  424. out.append(self.detect_and_crop(img))
  425. except Exception:
  426. out.append(None)
  427. return out