#!/usr/bin/env python3 # -*- coding: utf-8 -*- """ 卡牌图像识别/匹配 HTTP 接口(后端规范契约 + 双区级联匹配) 用途: 后端传图片 URL → 返回 Top-K card_id 及 match_rate。 匹配链路与入库完全体 / serve_card_match.py / scripts/cascade_match.py 同源: 下载图 → YOLO 裁卡 → letterbox 392 → 上半区 DINOv2-large 全库召回 Top-K_recall(物种级) → 下半区 Layer3 对候选算版本相似度 → fusion = α·upper + (1-α)·lower 重排 → card_id 请求/响应契约不变(code/message/data.matches...)。 match_rate = round(fusion × 100),0-100。 model_version = card_seg_v0904 OCR sidecar(serve_lang_judge.py :8100) 仍可选:仅记 language 到日志, 不参与检索前置(与入库双区链路一致:语种消歧在入库后处理,不在本接口做分层检索)。 环境:pytorch conda。 启动: CUDA_VISIBLE_DEVICES= \ ~/miniconda3/envs/pytorch/bin/python serve_recognition_api.py --port 8010 """ import os import sys import time import traceback import tempfile import csv sys.stdout.reconfigure(encoding="utf-8") sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) _SCRIPTS = os.path.join(os.path.dirname(os.path.abspath(__file__)), "scripts") if _SCRIPTS not in sys.path: sys.path.insert(0, _SCRIPTS) import numpy as np from PIL import Image import config from modules import image_downloader from cascade_match import DualCascadeMatcher # ==================== 服务级常量 ==================== MODEL_VERSION = "card_seg_v0904" # 版本号带图库版次后缀(v819=gallery_v819,2026-08-19 版库/权重) TOP_K_MAX = 20 LANG_JUDGE_URL = os.environ.get("LANG_JUDGE_URL", "http://127.0.0.1:8100/lang_judge") LANG_JUDGE_TIMEOUT = float(os.environ.get("LANG_JUDGE_TIMEOUT", "8")) API_TOKEN = os.environ.get("API_TOKEN") # ==================== 启动时一次性加载 ==================== print("[启动] 加载双区级联匹配器(上半区召回 + 下半区重排)...", flush=True) MATCHER = DualCascadeMatcher( alpha=config.CASCADE_ALPHA, top_k_recall=config.CASCADE_TOP_K_RECALL, verbose=True, ) print("[启动] 图库大小:", len(MATCHER.card_ids), "α=", MATCHER.alpha, "K_recall=", MATCHER.top_k_recall, flush=True) CARD_FIELDS = {} _FIELDS_CSV = os.path.join(config.DATA_DIR, "card_master_fields.csv") if os.path.exists(_FIELDS_CSV): with open(_FIELDS_CSV, encoding="utf-8-sig") as _f: for _r in csv.DictReader(_f): CARD_FIELDS[_r["card_id"]] = { "card_name": _r.get("card_name"), "card_name_ch": _r.get("card_name_ch"), "pg_label": _r.get("pg_label"), } print("[启动] 卡牌字段表: %d 张" % len(CARD_FIELDS), flush=True) else: print("[启动] ⚠ 未找到 card_master_fields.csv,algorithm_* 回退 dual meta", flush=True) _ID2META = {} for _m in (MATCHER.metas or []): cid = (_m or {}).get("card_id") if cid and cid not in _ID2META: _ID2META[cid] = _m print("[启动] 就绪,model_version=", MODEL_VERSION, flush=True) # ==================== 内部工具 ==================== class RecognizeError(Exception): def __init__(self, error_code, http_code, message): super().__init__(message) self.error_code = error_code self.http_code = http_code def _names(card_id): row = CARD_FIELDS.get(card_id) if row: name = row.get("card_name") or row.get("card_name_ch") return name, row.get("pg_label") m = _ID2META.get(card_id) or {} return (m.get("card_name") or m.get("card_name_ch")), m.get("pg_label") def recognize(image_url, top_k=None, min_match_rate=None, return_embedding=False): """图片 URL → matches(双区级联)。""" top_k = config.TOP_K if top_k is None else max(1, min(int(top_k), TOP_K_MAX)) min_match_rate = 0 if min_match_rate is None else int(min_match_rate) config.ensure_dirs() img_path = image_downloader.download_card_image(image_url, config.QUERY_IMG_DIR) if not img_path: raise RecognizeError("INVALID_IMAGE_URL", 400, "image url is not accessible: " + image_url) to_delete = img_path try: # 快速校验可读 try: Image.open(img_path).convert("RGB") except Exception: raise RecognizeError("UNSUPPORTED_IMAGE_TYPE", 415, "unsupported or corrupted image content") # 双区级联:与 cascade_match / 入库完全体同源 result_q = MATCHER.query(img_path, top_n=top_k) if result_q is None: raise RecognizeError("NO_CARD_DETECTED", 422, "no card detected in image or feature extraction failed") matches = [] for item in (result_q.get("top_k") or []): fusion = float(item.get("fusion", 0)) match_rate = int(max(0, min(100, round(fusion * 100)))) if match_rate < min_match_rate: continue cid = item["card_id"] name, series = _names(cid) matches.append({ "algorithm_card_name": name, "algorithm_series_name": series, "card_id": cid, "match_rate": match_rate, "rank": len(matches) + 1, }) data = { "matches": matches, "_fusion": float(result_q.get("fusion_score") or 0), "_upper": float(result_q.get("upper_sim") or 0), "_lower": float(result_q.get("lower_sim") or 0), } # return_embedding:返回上半区 1024 维(双区无单一 768 维向量) if return_embedding: try: q_u, q_l = MATCHER._features_of(img_path) if q_u is not None: data["embedding"] = q_u.astype(float).tolist() data["embedding_lower"] = q_l.astype(float).tolist() data["embedding_dim"] = 1024 except Exception: pass return data finally: if to_delete and os.path.exists(to_delete): try: os.remove(to_delete) except Exception: pass def _warmup(): """假图写临时文件,跑一次双区链路,编译 CUDA kernel。""" try: t0 = time.perf_counter() dummy = np.random.randint(40, 220, (640, 448, 3), dtype=np.uint8) fd, tmp = tempfile.mkstemp(suffix=".jpg") os.close(fd) try: Image.fromarray(dummy).save(tmp, quality=90) MATCHER.query(tmp, top_n=config.TOP_K) finally: if os.path.exists(tmp): try: os.remove(tmp) except Exception: pass print("[启动] 预热完成 (%.0f ms)" % ((time.perf_counter() - t0) * 1000), flush=True) except Exception as e: print("[启动] 预热跳过:", str(e)[:160], flush=True) _warmup() # ==================== HTTP ==================== def create_app(): from flask import Flask, request, jsonify app = Flask(__name__) app.json.sort_keys = False app.json.ensure_ascii = False @app.get("/health") def health(): return jsonify({ "status": "ok", "gallery_size": len(MATCHER.card_ids), "model_version": MODEL_VERSION, "alpha": MATCHER.alpha, "top_k_recall": MATCHER.top_k_recall, "feature_dim": 1024, }) def _check_auth(): if not API_TOKEN: return True auth = request.headers.get("Authorization", "") return auth == "Bearer " + API_TOKEN @app.post("/recognize") def recognize_endpoint(): t0 = time.perf_counter() req_id = request.headers.get("X-Request-Id") if not _check_auth(): return jsonify({"code": 401, "message": "unauthorized", "data": None}), 401 body = request.get_json(silent=True) or {} task_id = body.get("task_id") image_url = body.get("image_url") if not task_id or not image_url: msg = "missing required field task_id or image_url" return jsonify({ "code": 400, "message": msg, "data": {"task_id": task_id or "", "status": "failed", "error_code": "INVALID_REQUEST", "error_message": msg}, }), 400 top_k = body.get("top_k") min_match_rate = body.get("min_match_rate") return_embedding = bool(body.get("return_embedding", False)) print("[recognize] task_id=%s x_req=%s url=%s" % (task_id, req_id, image_url), flush=True) try: data = recognize(image_url, top_k=top_k, min_match_rate=min_match_rate, return_embedding=return_embedding) elapsed = int(round((time.perf_counter() - t0) * 1000)) resp_data = { "task_id": task_id, "status": "success", "model_version": MODEL_VERSION, "processing_time_ms": elapsed, "matches": data["matches"], } if "embedding" in data: resp_data["embedding"] = data["embedding"] if "embedding_lower" in data: resp_data["embedding_lower"] = data["embedding_lower"] return jsonify({"code": 200, "message": "success", "data": resp_data}) except RecognizeError as re: return jsonify({ "code": re.http_code, "message": str(re), "data": {"task_id": task_id, "status": "failed", "error_code": re.error_code, "error_message": str(re)}, }), re.http_code except Exception as e: traceback.print_exc() msg = str(e) or "unknown model error" return jsonify({ "code": 500, "message": msg, "data": {"task_id": task_id, "status": "failed", "error_code": "MODEL_INFERENCE_ERROR", "error_message": msg}, }), 500 return app if __name__ == "__main__": import argparse p = argparse.ArgumentParser(description="卡牌图像识别/匹配 HTTP 接口(双区级联)") p.add_argument("--host", default="0.0.0.0") p.add_argument("--port", type=int, default=8010) args = p.parse_args() create_app().run(host=args.host, port=args.port, threaded=True)