serve_recognition_api.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282
  1. #!/usr/bin/env python3
  2. # -*- coding: utf-8 -*-
  3. """
  4. 卡牌图像识别/匹配 HTTP 接口(后端规范契约 + 双区级联匹配)
  5. 用途:
  6. 后端传图片 URL → 返回 Top-K card_id 及 match_rate。
  7. 匹配链路与入库完全体 / serve_card_match.py / scripts/cascade_match.py 同源:
  8. 下载图 → YOLO 裁卡 → letterbox 392
  9. → 上半区 DINOv2-large 全库召回 Top-K_recall(物种级)
  10. → 下半区 Layer3 对候选算版本相似度
  11. → fusion = α·upper + (1-α)·lower 重排 → card_id
  12. 请求/响应契约不变(code/message/data.matches...)。
  13. match_rate = round(fusion × 100),0-100。
  14. model_version = card_seg_v0904
  15. OCR sidecar(serve_lang_judge.py :8100) 仍可选:仅记 language 到日志,
  16. 不参与检索前置(与入库双区链路一致:语种消歧在入库后处理,不在本接口做分层检索)。
  17. 环境:pytorch conda。
  18. 启动:
  19. CUDA_VISIBLE_DEVICES=<GPU UUID> \
  20. ~/miniconda3/envs/pytorch/bin/python serve_recognition_api.py --port 8010
  21. """
  22. import os
  23. import sys
  24. import time
  25. import traceback
  26. import tempfile
  27. import csv
  28. sys.stdout.reconfigure(encoding="utf-8")
  29. sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
  30. _SCRIPTS = os.path.join(os.path.dirname(os.path.abspath(__file__)), "scripts")
  31. if _SCRIPTS not in sys.path:
  32. sys.path.insert(0, _SCRIPTS)
  33. import numpy as np
  34. from PIL import Image
  35. import config
  36. from modules import image_downloader
  37. from cascade_match import DualCascadeMatcher
  38. # ==================== 服务级常量 ====================
  39. MODEL_VERSION = "card_seg_v0904" # 版本号带图库版次后缀(v819=gallery_v819,2026-08-19 版库/权重)
  40. TOP_K_MAX = 20
  41. LANG_JUDGE_URL = os.environ.get("LANG_JUDGE_URL", "http://127.0.0.1:8100/lang_judge")
  42. LANG_JUDGE_TIMEOUT = float(os.environ.get("LANG_JUDGE_TIMEOUT", "8"))
  43. API_TOKEN = os.environ.get("API_TOKEN")
  44. # ==================== 启动时一次性加载 ====================
  45. print("[启动] 加载双区级联匹配器(上半区召回 + 下半区重排)...", flush=True)
  46. MATCHER = DualCascadeMatcher(
  47. alpha=config.CASCADE_ALPHA,
  48. top_k_recall=config.CASCADE_TOP_K_RECALL,
  49. verbose=True,
  50. )
  51. print("[启动] 图库大小:", len(MATCHER.card_ids),
  52. "α=", MATCHER.alpha, "K_recall=", MATCHER.top_k_recall, flush=True)
  53. CARD_FIELDS = {}
  54. _FIELDS_CSV = os.path.join(config.DATA_DIR, "card_master_fields.csv")
  55. if os.path.exists(_FIELDS_CSV):
  56. with open(_FIELDS_CSV, encoding="utf-8-sig") as _f:
  57. for _r in csv.DictReader(_f):
  58. CARD_FIELDS[_r["card_id"]] = {
  59. "card_name": _r.get("card_name"),
  60. "card_name_ch": _r.get("card_name_ch"),
  61. "pg_label": _r.get("pg_label"),
  62. }
  63. print("[启动] 卡牌字段表: %d 张" % len(CARD_FIELDS), flush=True)
  64. else:
  65. print("[启动] ⚠ 未找到 card_master_fields.csv,algorithm_* 回退 dual meta", flush=True)
  66. _ID2META = {}
  67. for _m in (MATCHER.metas or []):
  68. cid = (_m or {}).get("card_id")
  69. if cid and cid not in _ID2META:
  70. _ID2META[cid] = _m
  71. print("[启动] 就绪,model_version=", MODEL_VERSION, flush=True)
  72. # ==================== 内部工具 ====================
  73. class RecognizeError(Exception):
  74. def __init__(self, error_code, http_code, message):
  75. super().__init__(message)
  76. self.error_code = error_code
  77. self.http_code = http_code
  78. def _names(card_id):
  79. row = CARD_FIELDS.get(card_id)
  80. if row:
  81. name = row.get("card_name") or row.get("card_name_ch")
  82. return name, row.get("pg_label")
  83. m = _ID2META.get(card_id) or {}
  84. return (m.get("card_name") or m.get("card_name_ch")), m.get("pg_label")
  85. def recognize(image_url, top_k=None, min_match_rate=None, return_embedding=False):
  86. """图片 URL → matches(双区级联)。"""
  87. top_k = config.TOP_K if top_k is None else max(1, min(int(top_k), TOP_K_MAX))
  88. min_match_rate = 0 if min_match_rate is None else int(min_match_rate)
  89. config.ensure_dirs()
  90. img_path = image_downloader.download_card_image(image_url, config.QUERY_IMG_DIR)
  91. if not img_path:
  92. raise RecognizeError("INVALID_IMAGE_URL", 400,
  93. "image url is not accessible: " + image_url)
  94. to_delete = img_path
  95. try:
  96. # 快速校验可读
  97. try:
  98. Image.open(img_path).convert("RGB")
  99. except Exception:
  100. raise RecognizeError("UNSUPPORTED_IMAGE_TYPE", 415,
  101. "unsupported or corrupted image content")
  102. # 双区级联:与 cascade_match / 入库完全体同源
  103. result_q = MATCHER.query(img_path, top_n=top_k)
  104. if result_q is None:
  105. raise RecognizeError("NO_CARD_DETECTED", 422,
  106. "no card detected in image or feature extraction failed")
  107. matches = []
  108. for item in (result_q.get("top_k") or []):
  109. fusion = float(item.get("fusion", 0))
  110. match_rate = int(max(0, min(100, round(fusion * 100))))
  111. if match_rate < min_match_rate:
  112. continue
  113. cid = item["card_id"]
  114. name, series = _names(cid)
  115. matches.append({
  116. "algorithm_card_name": name,
  117. "algorithm_series_name": series,
  118. "card_id": cid,
  119. "match_rate": match_rate,
  120. "rank": len(matches) + 1,
  121. })
  122. data = {
  123. "matches": matches,
  124. "_fusion": float(result_q.get("fusion_score") or 0),
  125. "_upper": float(result_q.get("upper_sim") or 0),
  126. "_lower": float(result_q.get("lower_sim") or 0),
  127. }
  128. # return_embedding:返回上半区 1024 维(双区无单一 768 维向量)
  129. if return_embedding:
  130. try:
  131. q_u, q_l = MATCHER._features_of(img_path)
  132. if q_u is not None:
  133. data["embedding"] = q_u.astype(float).tolist()
  134. data["embedding_lower"] = q_l.astype(float).tolist()
  135. data["embedding_dim"] = 1024
  136. except Exception:
  137. pass
  138. return data
  139. finally:
  140. if to_delete and os.path.exists(to_delete):
  141. try:
  142. os.remove(to_delete)
  143. except Exception:
  144. pass
  145. def _warmup():
  146. """假图写临时文件,跑一次双区链路,编译 CUDA kernel。"""
  147. try:
  148. t0 = time.perf_counter()
  149. dummy = np.random.randint(40, 220, (640, 448, 3), dtype=np.uint8)
  150. fd, tmp = tempfile.mkstemp(suffix=".jpg")
  151. os.close(fd)
  152. try:
  153. Image.fromarray(dummy).save(tmp, quality=90)
  154. MATCHER.query(tmp, top_n=config.TOP_K)
  155. finally:
  156. if os.path.exists(tmp):
  157. try:
  158. os.remove(tmp)
  159. except Exception:
  160. pass
  161. print("[启动] 预热完成 (%.0f ms)" % ((time.perf_counter() - t0) * 1000), flush=True)
  162. except Exception as e:
  163. print("[启动] 预热跳过:", str(e)[:160], flush=True)
  164. _warmup()
  165. # ==================== HTTP ====================
  166. def create_app():
  167. from flask import Flask, request, jsonify
  168. app = Flask(__name__)
  169. app.json.sort_keys = False
  170. app.json.ensure_ascii = False
  171. @app.get("/health")
  172. def health():
  173. return jsonify({
  174. "status": "ok",
  175. "gallery_size": len(MATCHER.card_ids),
  176. "model_version": MODEL_VERSION,
  177. "alpha": MATCHER.alpha,
  178. "top_k_recall": MATCHER.top_k_recall,
  179. "feature_dim": 1024,
  180. })
  181. def _check_auth():
  182. if not API_TOKEN:
  183. return True
  184. auth = request.headers.get("Authorization", "")
  185. return auth == "Bearer " + API_TOKEN
  186. @app.post("/recognize")
  187. def recognize_endpoint():
  188. t0 = time.perf_counter()
  189. req_id = request.headers.get("X-Request-Id")
  190. if not _check_auth():
  191. return jsonify({"code": 401, "message": "unauthorized", "data": None}), 401
  192. body = request.get_json(silent=True) or {}
  193. task_id = body.get("task_id")
  194. image_url = body.get("image_url")
  195. if not task_id or not image_url:
  196. msg = "missing required field task_id or image_url"
  197. return jsonify({
  198. "code": 400, "message": msg,
  199. "data": {"task_id": task_id or "", "status": "failed",
  200. "error_code": "INVALID_REQUEST", "error_message": msg},
  201. }), 400
  202. top_k = body.get("top_k")
  203. min_match_rate = body.get("min_match_rate")
  204. return_embedding = bool(body.get("return_embedding", False))
  205. print("[recognize] task_id=%s x_req=%s url=%s" % (task_id, req_id, image_url), flush=True)
  206. try:
  207. data = recognize(image_url, top_k=top_k, min_match_rate=min_match_rate,
  208. return_embedding=return_embedding)
  209. elapsed = int(round((time.perf_counter() - t0) * 1000))
  210. resp_data = {
  211. "task_id": task_id,
  212. "status": "success",
  213. "model_version": MODEL_VERSION,
  214. "processing_time_ms": elapsed,
  215. "matches": data["matches"],
  216. }
  217. if "embedding" in data:
  218. resp_data["embedding"] = data["embedding"]
  219. if "embedding_lower" in data:
  220. resp_data["embedding_lower"] = data["embedding_lower"]
  221. return jsonify({"code": 200, "message": "success", "data": resp_data})
  222. except RecognizeError as re:
  223. return jsonify({
  224. "code": re.http_code, "message": str(re),
  225. "data": {"task_id": task_id, "status": "failed",
  226. "error_code": re.error_code, "error_message": str(re)},
  227. }), re.http_code
  228. except Exception as e:
  229. traceback.print_exc()
  230. msg = str(e) or "unknown model error"
  231. return jsonify({
  232. "code": 500, "message": msg,
  233. "data": {"task_id": task_id, "status": "failed",
  234. "error_code": "MODEL_INFERENCE_ERROR", "error_message": msg},
  235. }), 500
  236. return app
  237. if __name__ == "__main__":
  238. import argparse
  239. p = argparse.ArgumentParser(description="卡牌图像识别/匹配 HTTP 接口(双区级联)")
  240. p.add_argument("--host", default="0.0.0.0")
  241. p.add_argument("--port", type=int, default=8010)
  242. args = p.parse_args()
  243. create_app().run(host=args.host, port=args.port, threaded=True)