| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475 |
- # -*- coding: utf-8 -*-
- """
- YOLO 卡牌检测裁切模块
- 默认查询路径(2026-08):card_seg_v2 (yolo26s-seg) → mask 4 点 → warp_perspective → ensure_portrait
- - 向量库/图库特征不动;仅优化「待匹配拍摄图」的裁切与几何摆正
- - TRT:CARD_SEG_USE_TRT=1 时优先加载同目录同名 .engine(.pt/.onnx 均可推导)
- - 回退:CARD_CROP_MODE=bbox_legacy 或 CARD_CROP_MODE 未设且仍指向旧 onnx 时走矩形 bbox 裁切
- 环境变量:
- CARD_SEG_USE_TRT=1 启用 TensorRT engine(须同目录存在 .engine)
- CARD_CROP_MODE=seg_warp|bbox_legacy 覆盖 config.CARD_CROP_MODE
- CARD_MERGE_DIST_PX 近距 mask 合并距离阈值(px),<=0 关闭;默认取 config(50)
- CARD_MERGE_MIN_IOU 合并可选 IoU 门槛,0=纯距离口径;默认取 config(0)
- """
- import math
- import os
- import numpy as np
- import cv2
- from modules.ultralytics_compat import import_yolo
- def _env_truthy(name, default="0"):
- return os.environ.get(name, default).lower() in ("1", "true", "yes", "seg_warp")
- def _resolve_engine_path(model_path):
- """同目录同名 .engine:xxx.pt / xxx.onnx → xxx.engine"""
- if model_path.endswith(".onnx"):
- return model_path[:-5] + ".engine"
- if model_path.endswith(".pt"):
- return model_path[:-3] + ".engine"
- if model_path.endswith(".engine"):
- return model_path
- return None
- def _pack_crop(crop, box=None, label=None, return_box=False, return_label=False):
- """按需打包 detect_and_crop 返回值:crop / (crop,box) / (crop,label) / (crop,box,label)。"""
- if return_label and return_box:
- return crop, box, label
- if return_label:
- return crop, label
- if return_box:
- return crop, box
- return crop
- def _sort_reading_order(items):
- """多卡按阅读顺序:先上后下分行,同行从左到右。
- items 每项需含 box=(x1,y1,x2,y2)。"""
- if not items:
- return items
- heights = [max(1.0, it["box"][3] - it["box"][1]) for it in items]
- row_tol = 0.5 * float(np.median(heights))
- ordered = sorted(
- items,
- key=lambda it: ((it["box"][1] + it["box"][3]) / 2.0,
- (it["box"][0] + it["box"][2]) / 2.0),
- )
- rows = []
- for it in ordered:
- cy = (it["box"][1] + it["box"][3]) / 2.0
- if rows and abs(cy - rows[-1]["cy"]) <= row_tol:
- rows[-1]["items"].append(it)
- n = len(rows[-1]["items"])
- rows[-1]["cy"] = (rows[-1]["cy"] * (n - 1) + cy) / n
- else:
- rows.append({"cy": cy, "items": [it]})
- out = []
- for row in rows:
- out.extend(sorted(row["items"],
- key=lambda it: (it["box"][0] + it["box"][2]) / 2.0))
- return out
- def _box_iou(a, b):
- """xyxy 框 IoU(近距合并的可选门槛用)。"""
- ix1, iy1 = max(a[0], b[0]), max(a[1], b[1])
- ix2, iy2 = min(a[2], b[2]), min(a[3], b[3])
- iw, ih = max(0.0, ix2 - ix1), max(0.0, iy2 - iy1)
- inter = iw * ih
- if inter <= 0:
- return 0.0
- area_a = max(0.0, a[2] - a[0]) * max(0.0, a[3] - a[1])
- area_b = max(0.0, b[2] - b[0]) * max(0.0, b[3] - b[1])
- union = area_a + area_b - inter
- return inter / union if union > 0 else 0.0
- def _center_of(mask, box, min_area_frac=0.01, stride=4):
- """实例中心:mask 质心(有效面积 >= 1% 图幅,防噪声拉偏),否则 bbox 中心兜底。
- 均为原图分辨率坐标 (cx, cy)。质心用 stride 跨步采样:retina mask 是原图分辨率
- (4256×2128 全图 nonzero 实测 ~10ms/实例,占 8020 post 段大头),采样后
- ±stride px 抖动远小于 50px 合并阈值,合并/排序行为不变。"""
- if mask is not None:
- sub = mask[::stride, ::stride] > 0.5
- ys, xs = np.nonzero(sub)
- if xs.size >= min_area_frac * sub.shape[0] * sub.shape[1]:
- return (float(xs.mean()) * stride, float(ys.mean()) * stride)
- return ((box[0] + box[2]) / 2.0, (box[1] + box[3]) / 2.0)
- class CardDetector:
- """卡牌区域检测 + 查询侧裁切(bbox 或 seg+warp)"""
- def __init__(self, model_path, conf=0.25, imgsz=640, device=None, crop_mode=None):
- YOLO = import_yolo()
- if not os.path.exists(model_path):
- raise FileNotFoundError(f"YOLO 模型不存在: {model_path}")
- self.conf = conf
- self.imgsz = imgsz
- self.model_path = model_path
- # crop_mode: seg_warp | bbox_legacy
- env_mode = os.environ.get("CARD_CROP_MODE", "").strip().lower()
- if crop_mode is not None:
- self.crop_mode = crop_mode
- elif env_mode in ("seg_warp", "bbox_legacy"):
- self.crop_mode = env_mode
- else:
- try:
- import config as _cfg
- self.crop_mode = getattr(_cfg, "CARD_CROP_MODE", "seg_warp")
- except Exception:
- self.crop_mode = "seg_warp"
- if self.crop_mode not in ("seg_warp", "bbox_legacy"):
- self.crop_mode = "seg_warp"
- # 近距 mask 合并阈值:env → config → 默认(detect_and_crop_all 用;<=0 关闭合并)
- def _merge_cfg(key, default):
- v = os.environ.get(key, "").strip()
- if v:
- try:
- return float(v)
- except ValueError:
- pass
- try:
- import config as _cfg
- return float(getattr(_cfg, key, default))
- except Exception:
- return float(default)
- self.merge_dist_px = _merge_cfg("CARD_MERGE_DIST_PX", 50.0)
- self.merge_min_iou = _merge_cfg("CARD_MERGE_MIN_IOU", 0.0)
- self._trt = False
- use_trt = os.environ.get("CARD_SEG_USE_TRT", "0").lower() in ("1", "true", "yes")
- engine_path = _resolve_engine_path(model_path) if use_trt else None
- # 若直接传入 .engine
- if model_path.endswith(".engine") and os.path.exists(model_path):
- engine_path = model_path
- use_trt = True
- if use_trt and engine_path and os.path.exists(engine_path):
- try:
- import tensorrt # noqa: F401
- self.model = YOLO(engine_path, task="segment")
- self.device = device if device is not None else 0
- self._trt = True
- print(
- f"[CardDetector] TRT engine mode={self.crop_mode} device={self.device} "
- f"engine={os.path.basename(engine_path)}",
- flush=True,
- )
- except Exception as e:
- print(
- f"[CardDetector] TRT 加载失败({type(e).__name__}: {e}), 回退 {model_path}",
- flush=True,
- )
- self.model = YOLO(model_path, task="segment")
- self.device = device
- else:
- self.model = YOLO(model_path, task="segment")
- self.device = device
- if use_trt and engine_path and not os.path.exists(engine_path):
- print(
- f"[CardDetector] CARD_SEG_USE_TRT=1 但 engine 不存在: {engine_path}, 用 {model_path}",
- flush=True,
- )
- def _predict(self, img):
- is_path = isinstance(img, str)
- results = self.model.predict(
- img if is_path else img,
- conf=self.conf,
- imgsz=self.imgsz,
- device=self.device,
- verbose=False,
- retina_masks=True,
- )
- return results, is_path
- @staticmethod
- def _top_label(res, idx):
- """取第 idx 个实例的类别名(yolo26s-seg 的 label,如 'pokemon'/'nba');取不到返回 None。"""
- try:
- boxes = res.boxes
- if boxes is None or boxes.cls is None or len(boxes.cls) <= idx:
- return None
- cls_id = int(boxes.cls[idx].item() if hasattr(boxes.cls[idx], "item") else boxes.cls[idx])
- return (res.names or {}).get(cls_id)
- except Exception:
- return None
- def detect_and_crop(self, img, return_box=False, return_label=False):
- """
- 检测并裁切最大卡牌区域。
- - crop_mode=bbox_legacy: 仅 xyxy 矩形裁切(旧行为)
- - crop_mode=seg_warp: mask→4点→warp(valid_quad 则始终 warp)→ensure_portrait
- 返回 RGB;失败时行为与旧版一致(尽量返回原图 RGB,return_box 时 box 可能为 None)
- return_label=True 时额外返回最大实例的 yolo label 名称(如 'pokemon'):
- (crop, label) 或 (crop, box, label);未检出卡时 label=None
- """
- if self.crop_mode == "seg_warp":
- return self._detect_and_crop_seg_warp(img, return_box=return_box, return_label=return_label)
- return self._detect_and_crop_bbox(img, return_box=return_box, return_label=return_label)
- def detect_and_crop_all(self, img):
- """检测图中全部卡牌实例,按阅读顺序(上→下、同行左→右)裁切。
- YOLO 只跑一次。近距合并:两实例中心(mask 质心,回退 bbox 中心)距离
- <= merge_dist_px 视为同一张卡的重复分割,仅保留 conf 高者(8020 / 8000
- 多卡路径共用本逻辑;阈值 CARD_MERGE_DIST_PX 可 env/config 调,<=0 关闭)。
- 返回 list[{"crop": RGB, "box": (x1,y1,x2,y2), "label": str|None,
- "conf": float, "center": (cx,cy)}],conf/center 为新增键(加法,
- query_all 等既有调用方不受影响)。
- 未检出返回 [](与 detect_and_crop 单卡失败回原图的行为刻意不同——多卡路径不能把整图当一张卡)。"""
- is_path = isinstance(img, str)
- results, _ = self._predict(img)
- if not results:
- return []
- res = results[0]
- orig = res.orig_img
- if orig is None:
- return []
- h, w = orig.shape[:2]
- orig_rgb = cv2.cvtColor(orig, cv2.COLOR_BGR2RGB) if is_path else orig
- boxes = res.boxes
- if boxes is None or len(boxes) == 0:
- return []
- xyxy = boxes.xyxy.cpu().numpy()
- confs = None
- try:
- if boxes.conf is not None and len(boxes.conf) == len(xyxy):
- confs = boxes.conf.cpu().numpy()
- except Exception:
- confs = None
- masks_np = None
- if res.masks is not None:
- masks_np = res.masks.data.cpu().numpy()
- # 第一段:逐实例元信息(跨步采样质心 + conf;全分辨率二值化推迟到幸存者,
- # 被吸收实例不再白做二值化——post 段 2026-09-03 profile 占单请求 ~11%)
- metas = []
- for i in range(len(xyxy)):
- box = tuple(float(v) for v in xyxy[i])
- mask_raw = None
- if masks_np is not None and i < len(masks_np):
- mask_raw = masks_np[i]
- if mask_raw.shape[:2] != (h, w):
- mask_raw = cv2.resize(mask_raw, (w, h), interpolation=cv2.INTER_NEAREST)
- metas.append({
- "box": box,
- "label": self._top_label(res, i),
- "mask_raw": mask_raw,
- "conf": float(confs[i]) if confs is not None else 0.0,
- "center": _center_of(mask_raw, box),
- })
- # 第二段:近距合并(同一张卡被拆成多个 mask / 重复检测 → 保 conf 高者)
- keep, _dropped = self._merge_indices(metas)
- # 第三段:只对幸存者二值化全分辨率 mask + 裁卡(被吸收者连二值化都不做)
- items = []
- for m in keep:
- mask = None
- if m["mask_raw"] is not None:
- mask = (m["mask_raw"] > 0.5).astype(np.uint8) * 255
- crop = self._crop_one_instance(orig_rgb, m["box"], mask)
- items.append({
- "crop": crop,
- "box": tuple(int(round(v)) for v in m["box"]),
- "label": m["label"],
- "conf": m["conf"],
- "center": m["center"],
- })
- return _sort_reading_order(items)
- def _merge_indices(self, metas):
- """近距合并:conf 降序贪心——每个实例与已保留实例算中心欧氏距,
- 距离 <= merge_dist_px(且 IoU >= merge_min_iou,若门槛>0 启用)则被吸收。
- 返回 (保留 metas 列表(按原顺序), 被吸收索引列表)。"""
- thr = self.merge_dist_px
- min_iou = self.merge_min_iou
- if thr <= 0 or len(metas) <= 1:
- return metas, []
- order = sorted(range(len(metas)), key=lambda i: -metas[i]["conf"])
- kept, dropped = [], []
- for i in order:
- mi = metas[i]
- dup = False
- for j in kept:
- mj = metas[j]
- d = math.hypot(mi["center"][0] - mj["center"][0],
- mi["center"][1] - mj["center"][1])
- if d > thr:
- continue
- if min_iou > 0 and _box_iou(mi["box"], mj["box"]) < min_iou:
- continue
- dup = True
- break
- if dup:
- dropped.append(i)
- else:
- kept.append(i)
- if dropped:
- print(f"[CardDetector] 近距合并(thr={thr:g}px"
- + (f" iou>={min_iou:g}" if min_iou > 0 else "")
- + f"): 保留 {sorted(kept)} 吸收 {sorted(dropped)} "
- + str([(metas[i]["center"], round(metas[i]["conf"], 3)) for i in dropped]),
- flush=True)
- return [metas[i] for i in sorted(kept)], dropped
- def _crop_one_instance(self, orig_rgb, box, mask):
- """单实例 → 裁卡 RGB(seg_warp 优先,失败回退 bbox)。"""
- from modules.rectifier import (
- mask_to_corners, order_corners, valid_quad, warp_perspective,
- crop_bbox, ensure_portrait,
- )
- out = None
- if self.crop_mode == "seg_warp" and mask is not None:
- corners = mask_to_corners(mask)
- if corners is not None:
- ordered = order_corners(corners)
- if valid_quad(ordered):
- warped = warp_perspective(orig_rgb, corners)
- if warped is not None and warped.size > 0:
- out = warped
- if out is None and box is not None:
- out = crop_bbox(orig_rgb, box)
- if out is None:
- out = orig_rgb
- out = ensure_portrait(out)
- if out.dtype != np.uint8:
- out = np.clip(out, 0, 255).astype(np.uint8)
- return out
- def _detect_and_crop_bbox(self, img, return_box=False, return_label=False):
- is_path = isinstance(img, str)
- results, _ = self._predict(img)
- if not results:
- return _pack_crop(None, return_box=return_box, return_label=return_label)
- res = results[0]
- orig = res.orig_img # 路径输入多为 BGR;数组输入约定与 ultralytics 一致
- if orig is None:
- return _pack_crop(None, return_box=return_box, return_label=return_label)
- boxes = res.boxes
- if boxes is None or len(boxes) == 0:
- rgb = cv2.cvtColor(orig, cv2.COLOR_BGR2RGB) if is_path else orig
- return _pack_crop(rgb, return_box=return_box, return_label=return_label)
- xyxy = boxes.xyxy.cpu().numpy()
- areas = (xyxy[:, 2] - xyxy[:, 0]) * (xyxy[:, 3] - xyxy[:, 1])
- max_idx = int(np.argmax(areas))
- label = self._top_label(res, max_idx)
- x1, y1, x2, y2 = xyxy[max_idx]
- h, w = orig.shape[:2]
- x1, y1 = max(0, int(x1)), max(0, int(y1))
- x2, y2 = min(w, int(x2)), min(h, int(y2))
- if x2 <= x1 or y2 <= y1:
- rgb = cv2.cvtColor(orig, cv2.COLOR_BGR2RGB) if is_path else orig
- return _pack_crop(rgb, return_box=return_box, return_label=return_label)
- crop = orig[y1:y2, x1:x2]
- crop_rgb = cv2.cvtColor(crop, cv2.COLOR_BGR2RGB) if is_path else crop
- return _pack_crop(crop_rgb, box=(x1, y1, x2, y2), label=label,
- return_box=return_box, return_label=return_label)
- def _detect_and_crop_seg_warp(self, img, return_box=False, return_label=False):
- """seg → 最大实例 mask → 4 点 → warp(始终尝试)→ ensure_portrait → RGB"""
- from modules.rectifier import (
- mask_to_corners, order_corners, valid_quad, warp_perspective,
- crop_bbox, ensure_portrait,
- )
- is_path = isinstance(img, str)
- if return_label:
- box, mask, orig_rgb, label = self.detect_box_mask(img, return_label=True)
- else:
- box, mask, orig_rgb = self.detect_box_mask(img)
- label = None
- if orig_rgb is None:
- return _pack_crop(None, return_box=return_box, return_label=return_label)
- # 在 RGB 上 warp,与下游 DINOv2 / letterbox 一致
- out = None
- if mask is not None:
- corners = mask_to_corners(mask)
- if corners is not None:
- ordered = order_corners(corners)
- # 查询侧:valid_quad 即 warp,不再用 has_perspective 跳过
- if valid_quad(ordered):
- warped = warp_perspective(orig_rgb, corners)
- if warped is not None and warped.size > 0:
- out = warped
- if out is None and box is not None:
- out = crop_bbox(orig_rgb, box)
- if out is None:
- out = orig_rgb
- out = ensure_portrait(out)
- # 保证 uint8 RGB
- if out.dtype != np.uint8:
- out = np.clip(out, 0, 255).astype(np.uint8)
- box_i = None
- if box is not None:
- box_i = tuple(int(round(v)) for v in box)
- return _pack_crop(out, box=box_i, label=label,
- return_box=return_box, return_label=return_label)
- def detect_box_mask(self, img, return_label=False):
- """
- 检测最大卡牌,返回 (box, mask, orig_rgb) 供透视矫正;return_label=True 时
- 返回 (box, mask, orig_rgb, label)(label 为 yolo 类别名,未检出时 None)。
- 颜色约定:路径输入 → RGB;数组输入原样返回(调用方保证语义)。
- """
- is_path = isinstance(img, str)
- results, _ = self._predict(img)
- if not results:
- return (None, None, None, None) if return_label else (None, None, None)
- res = results[0]
- orig = res.orig_img
- if orig is None:
- return (None, None, None, None) if return_label else (None, None, None)
- h, w = orig.shape[:2]
- # ultralytics 路径读图为 BGR;numpy 输入时 orig 通常与输入同通道序
- if is_path:
- orig_rgb = cv2.cvtColor(orig, cv2.COLOR_BGR2RGB)
- else:
- orig_rgb = orig
- boxes = res.boxes
- if boxes is None or len(boxes) == 0:
- return (None, None, orig_rgb, None) if return_label else (None, None, orig_rgb)
- xyxy = boxes.xyxy.cpu().numpy()
- areas = (xyxy[:, 2] - xyxy[:, 0]) * (xyxy[:, 3] - xyxy[:, 1])
- idx = int(np.argmax(areas))
- box = tuple(float(v) for v in xyxy[idx])
- label = self._top_label(res, idx)
- mask = None
- if res.masks is not None:
- m = res.masks.data.cpu().numpy()[idx]
- m = (m > 0.5).astype(np.uint8) * 255
- if m.shape[:2] != (h, w):
- m = cv2.resize(m, (w, h), interpolation=cv2.INTER_NEAREST)
- mask = m
- if return_label:
- return box, mask, orig_rgb, label
- return box, mask, orig_rgb
- def crop_batch(self, img_list):
- out = []
- for img in img_list:
- try:
- out.append(self.detect_and_crop(img))
- except Exception:
- out.append(None)
- return out
|