File size: 8,846 Bytes
3bda1a9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 | import os, json
import numpy as np
from pathlib import Path
from typing import Iterable, Optional, Union
from ultralytics import YOLO
import re
promotion_keywords = [
r"领", r"申请", r"立即",r"马上",r"即刻",r"立即",r"下载",
r"一定",r"现发",
r"大牌美食",
]
extra_keywords = [
r"立即领取",r"点击领取",r"申请我的额度",r"免费观看",
r"火热选购",
r"立即前往",r"立即下载",r"立即投保",r"立即参与",r"立即解救",r"立即领取",r"立即抢购",r"立即购买",r"立即签到",
r"测一测",r"一定要买",r"现摘现发",
r"登录",r"上滑",
]
compiled_patterns = [re.compile(p, flags=re.IGNORECASE) for p in (promotion_keywords + extra_keywords)]
compiled_extra_patterns = [re.compile(p, flags=re.IGNORECASE) for p in extra_keywords]
def extract_texts(ocr_result) -> str:
# 支持 str / [str] / dict / list[dict/...]
texts = []
texts = ocr_result[0]['rec_texts']
scores = ocr_result[0]['rec_scores']
coordss = ocr_result[0]['rec_polys']
ocr_result = [
(pts, (txt, conf))
for pts, txt, conf in zip(coordss, texts, scores)
]
for res in ocr_result:
coords, (text, confidence) = res
if confidence > 0.5:
texts.append(text)
merged = " ".join([t for t in texts if t]).strip()
merged = re.sub(r"\s+", "", merged)
print(merged)
return merged
def merge_by_rows(ocr_result, y_threshold=40):
blocks = []
texts = ocr_result[0]['rec_texts']
scores = ocr_result[0]['rec_scores']
coordss = ocr_result[0]['rec_polys']
ocr_result = [
(pts, (txt, conf))
for pts, txt, conf in zip(coordss, texts, scores)
]
for pts, (txt, conf) in ocr_result:
if conf < 0.5 or not txt.strip():
continue
ys = [p[1] for p in pts]
xs = [p[0] for p in pts]
blocks.append({
"text": txt.strip(),
"x_min": min(xs),
"y_max": max(ys)
})
if not blocks:
return ""
blocks.sort(key=lambda b: b["y_max"])
lines, cur = [], [blocks[0]]
for blk in blocks[1:]:
if abs(blk["y_max"] - cur[0]["y_max"]) <= y_threshold:
cur.append(blk)
else:
lines.append(cur)
cur = [blk]
lines.append(cur)
merged = []
for line in lines:
line.sort(key=lambda b: b["x_min"])
merged.append("".join(b["text"] for b in line))
return " ".join(merged)
# ========== 依赖函数(全部保留,原样引入即可) ==========
def is_guide_button_with_yolo_image(
image_np: np.ndarray,
model: Union[str, YOLO],
ocr,
out_dir: Union[str, Path] = "yolo_outputs",
name: str = "predict",
imgsz: int = 1024,
device: Union[int, str] = 1,
iou: float = 0.45,
max_det: int = 300,
half: bool = False,
verbose: bool = True,
idx: Optional[Union[int, str]] = None # 新增:指定保存文件名
) -> bool:
"""
用 YOLO 对单张 numpy 图像进行检测,并保存可视化与 JSON。
返回:是否包含促销元素(True/False)
"""
class_thresholds = {
0: 1.0, # 优惠券
1: 0.75, # 促销贴纸
2: 1.0, # 元宝
3: 1.0, # 大礼物盒子
4: 0.95, # 爆炸
5: 0.95, # 礼物盒子
6: 0.95, # 福袋
7: 0.95, # 红包
8: 0.95, # 金币
9: 0.95, # 金币簇
10: 0.35, # 长条促销贴纸
11: 0.35, # 长条贴纸
}
# 加载模型
if isinstance(model, str):
model = YOLO(model)
out_dir = Path(out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
results = model.predict(
source=[image_np],
imgsz=imgsz,
conf=0.25,
iou=iou,
device=device,
max_det=max_det,
half=half,
save=True,
project=str(out_dir),
name=name,
exist_ok=True,
verbose=verbose
)
promo_found = False
det_list = []
names_dict = getattr(model, "names", None)
for result in results:
save_dir = Path(result.save_dir)
if result.boxes is None:
continue
boxes = result.boxes.xyxy.cpu().numpy()
classes = result.boxes.cls.cpu().numpy().astype(int)
scores = result.boxes.conf.cpu().numpy()
if verbose:
print("—— detections ——")
h, w = image_np.shape[:2]
total_img_area = h * w # 图像总面积
for (x1, y1, x2, y2), cid, score in zip(boxes, classes, scores):
cname = names_dict.get(int(cid), str(int(cid))) if isinstance(names_dict, dict) else str(int(cid))
if verbose:
print(f"cls={cname}({cid}) conf={score:.3f}")
det_list.append([cname])
det_list.append([float(x1), float(y1), float(x2), float(y2), int(cid), float(score)])
th = class_thresholds.get(cid, None)
if th is None:
continue
# ================= 修改部分开始 =================
# 处理非长条类型的类别 (即 0-9)
if cid not in (10, 11):
if cid == 1:
# 计算当前检测框的面积
box_area = (x2 - x1) * (y2 - y1)
# 计算占比
ratio = box_area / total_img_area if total_img_area > 0 else 0
# 逻辑:如果是类别1,置信度达标 且 面积占比小于 0.01,判定为True
if score >= th and ratio < 0.01:
if verbose:
print(f" [Match] Class 1 (Area Ratio: {ratio:.4f} < 0.01)")
promo_found = True
else:
# 其他类别 (0, 2-9) 保持原逻辑:只看置信度
if score >= th:
promo_found = True
continue
# ================= 修改部分结束 =================
# 下面是处理类别 10, 11 的逻辑 (原代码保持不变)
if cid in (10,11) and score >= 0.875:
promo_found = True
continue
if score >= th:
if ocr is not None:
# 裁剪检测框区域
xi1, yi1 = max(0, int(x1)), max(0, int(y1))
xi2, yi2 = min(w - 1, int(x2)), min(h - 1, int(y2))
if xi2 > xi1 and yi2 > yi1:
crop = image_np[yi1:yi2, xi1:xi2]
try:
ocr_result = ocr.ocr(crop)
# text_merged = extract_texts(ocr_result)
text_merged = merge_by_rows(ocr_result)
print(text_merged)
# 正则匹配:命中任一关键词则判定为促销
if any(p.search(text_merged) for p in compiled_patterns):
if verbose:
print(f"[OCR-MATCH] {cname} 命中促销关键词")
promo_found = True
continue
except Exception as e:
if verbose:
print(f"[OCR-ERROR] {e}")
# OCR 失败时不改变 promo_found,保持谨慎
# —— 重命名保存的图片 ——
if idx is not None:
default_path = save_dir / "image0.jpg"
target_path = save_dir / f"{idx}.jpg"
if default_path.exists():
default_path.rename(target_path)
if verbose:
print(f"[保存] 可视化检测图: {target_path}")
# 保存检测结果 JSON
json_path = save_dir / f"detections_{idx if idx is not None else '0'}.json"
with open(json_path, "w", encoding="utf-8") as f:
json.dump(det_list, f, ensure_ascii=False, indent=2)
if verbose:
print(f"[保存] 检测数组: {json_path}")
# ==== 检测框全部没命中,再做一次全局OCR,匹配extra_keywords ====
if ocr is not None and not promo_found:
try:
global_ocr_result = ocr.ocr(image_np)
global_text = merge_by_rows(global_ocr_result)
print("[全局OCR]", global_text)
if any(p.search(global_text) for p in compiled_extra_patterns):
if verbose:
print("[全局OCR-MATCH] 命中extra_keywords促销关键词")
return True
except Exception as e:
if verbose:
print(f"[全局OCR-ERROR] {e}")
return promo_found |