| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282 |
- #!/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=<GPU UUID> \
- ~/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)
|