parse_grading_ocr.py 9.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242
  1. # -*- coding: utf-8 -*-
  2. """
  3. parse_grading_ocr.py - 从 best.pt 检测结果 + OCR 分行文本解析出结构化字段
  4. 字段:
  5. - detected_company: 来自 best.pt label 的前缀(如 BGS-AUTHENTIC -> BGS),多框取最高置信度
  6. - ocr_score: OCR 识别的评分(如 10, 9.5, 8)
  7. - score_desc: 评分描述(可选,如 GEM MT/NM-MT/MINT,在分数上方那行)
  8. - serial_no: 至少 6 个连续数字字符,无则 null
  9. """
  10. import csv
  11. import json
  12. import re
  13. from pathlib import Path
  14. ROOT = Path(r"D:/顾工交接/wzj/data_all/test/grading_ocr")
  15. OCR_JSON = ROOT / "ocr_result.json"
  16. OCR_CSV = ROOT / "ocr_summary.csv"
  17. import os
  18. # 若目标文件被占用,自动换名输出
  19. OUT_CSV = ROOT / "grading_parsed.csv"
  20. if OUT_CSV.exists():
  21. # 尝试用 _v2/_v3 之类,直到可用
  22. for i in range(2, 10):
  23. cand = ROOT / f"grading_parsed_v{i}.csv"
  24. if not cand.exists():
  25. OUT_CSV = cand
  26. print(f"⚠ grading_parsed.csv 被占用,改用 {OUT_CSV.name}")
  27. break
  28. # 读 OCR JSON
  29. with open(OCR_JSON, "r", encoding="utf-8") as f:
  30. ocr_data = json.load(f)
  31. # 读 CSV
  32. rows = []
  33. with open(OCR_CSV, "r", encoding="utf-8-sig") as f:
  34. for r in csv.DictReader(f):
  35. rows.append(r)
  36. def detect_company_from_label(box_label):
  37. """从 best.pt label(如 'BGS-AUTHENTIC','PSA-AUTO','SGC-OLD')提取公司前缀。
  38. Returns:
  39. ("PSA", 1) 或 ("BGS", 1) 等有框
  40. ("", 0) 无框(非评级卡)
  41. """
  42. if not box_label or box_label == "(无检测)":
  43. return "", 0
  44. return box_label.split("-", 1)[0], 1
  45. # score_desc 关键词 + 对应分数(用于校验,但 score_desc 是 OCR 文本本身,不强制对应)
  46. DESC_TO_SCORE = [
  47. ("GEM MT", 10), # PSA: Gem Mint = 10
  48. ("GEMMINT", 10),
  49. ("PRISTINE", 10),
  50. ("GEM", 10),
  51. ("MINT", 9), # PSA: Mint = 9
  52. ("NM-MT", 8.5),
  53. ("NMMT", 8.5),
  54. ("NM", 7), # 近 Mint
  55. ("EX", 6),
  56. ("EX+", 6.5),
  57. ("VG-EX", 4),
  58. ("VG", 4),
  59. ("GOOD", 3),
  60. ("FR", 1),
  61. ("PR", 1.5),
  62. ]
  63. def parse_score_and_desc(lines):
  64. """从 OCR 分行文本里提取:
  65. - score: 数字分数(1-10, 或 AU)
  66. - score_desc: 描述(GEM MT/NM-MT/MINT 等),必须在分数上方一行
  67. 关键改进:
  68. 1. 分数限定范围 1.0-10.0(排除年份/编号)
  69. 2. 描述必须在分数"紧邻的上一行"(不能跨行)
  70. 3. 描述词必须在专门的描述关键词表里(避免"PRIZMW/C"这种卡名误匹配)
  71. """
  72. score = None
  73. score_desc = ""
  74. score_line_idx = -1
  75. # 找分数行: 数字必须 1-10 范围内(年份/编号/日期/超大数都排除)
  76. for i, ln in enumerate(lines):
  77. ln_clean = ln.strip()
  78. # 匹配 1-10 范围内的纯数字(支持 .5)
  79. m = re.match(r"^(\d+(\.\d+)?)$", ln_clean)
  80. if m:
  81. val = float(m.group(1))
  82. if 1.0 <= val <= 10.0:
  83. score = ln_clean
  84. score_line_idx = i
  85. break
  86. # 单独 "AU" 也算分数(真伪)
  87. if ln_clean.upper() == "AU":
  88. score = "AU"
  89. score_line_idx = i
  90. break
  91. # 描述关键词白名单(只在分数"紧邻上一行"找,且必须是这些词)
  92. DESC_KEYWORDS = ["GEM MT", "GEMMINT", "GEM MINT", "PRISTINE",
  93. "MINT", "NM-MT", "NMMT", "NM",
  94. "EX", "EX+", "VG-EX", "VG", "GOOD", "PR", "FR",
  95. "AUTHENTIC"] # AUTHENTIC 是签字的真伪描述
  96. # 找分数紧邻上一行的描述
  97. if score_line_idx > 0:
  98. prev_line = lines[score_line_idx - 1].strip().upper()
  99. # 检查上一行是否纯是描述词
  100. for kw in DESC_KEYWORDS:
  101. if kw in prev_line or prev_line == kw.replace(" ", ""):
  102. score_desc = lines[score_line_idx - 1].strip()
  103. break
  104. # 兜底: 如果紧邻上行没找到,扫描所有行找独立的一行(整行就是描述词)
  105. if not score_desc:
  106. for ln in lines:
  107. ln_stripped = ln.strip().upper()
  108. if ln_stripped in [k.replace(" ", "") for k in DESC_KEYWORDS]:
  109. score_desc = ln.strip()
  110. break
  111. return score or "", score_desc
  112. def parse_serial(lines):
  113. """找至少 6 个连续数字字符(编号)。"""
  114. for ln in lines:
  115. m = re.search(r"\d{6,}", ln)
  116. if m:
  117. return m.group(0)
  118. return None # 无则 null
  119. # best.pt 的检测结果在 infer_test.py 生成的 test_predictions.json 里
  120. # 我们之前下载过 test_predict/test_predictions.json,但那是 infer_test 的输出
  121. # 这里用 OCR_CSV 里的 filename 作为 key,通过另一份 best.pt 检测数据来获取 label
  122. # 实际: best.pt 的 label 在 grading_ocr.py 第一步裁剪时记录了,但 crop_records.json 里只有 box/conf,没有 label
  123. # 解决方案: 重新跑一遍 best.pt 拿到 label,或从测试结果里取
  124. # 检查有没有 best.pt label 的数据
  125. # 方案: 通过 test_predictions.json (infer_test 生成的) 拿到 detected label
  126. PRED_JSON = Path(r"D:/顾工交接/wzj/data_all/test/test_predict/test_predictions.json")
  127. pred_data = {}
  128. if PRED_JSON.exists():
  129. with open(PRED_JSON, "r", encoding="utf-8") as f:
  130. pred_data = json.load(f)
  131. # pred_data 格式: {"records": [{true, pred, conf, file}, ...]}
  132. records = pred_data.get("records", [])
  133. pred_label = {r["file"]: r["pred"] for r in records}
  134. print(f"✓ 从 test_predictions.json 加载 best.pt 预测结果: {len(pred_label)} 条")
  135. else:
  136. pred_label = {}
  137. print("⚠ test_predictions.json 不存在,detected_company 字段将为空")
  138. # 解析
  139. out_rows = []
  140. for r in rows:
  141. name = r["filename"]
  142. lines = ocr_data.get(name, {}).get("lines", [])
  143. box_label = pred_label.get(name, "")
  144. detected_company, is_graded_int = detect_company_from_label(box_label)
  145. score, score_desc = parse_score_and_desc(lines)
  146. serial_no = parse_serial(lines)
  147. # 评级卡但 OCR 没识别出分数 -> 填 "authentic"(说明卡是真的评分过,但具体分数OCR漏识)
  148. if is_graded_int == 1 and not score:
  149. score = "authentic"
  150. # 非评级卡 -> 一律清空(OCR在原图上抓到的数字是噪音,不是评级分数)
  151. if is_graded_int == 0:
  152. score = ""
  153. out_rows.append({
  154. "filename": name,
  155. "true_label": r["true_label"],
  156. "detected_company": detected_company, # 来自 best.pt label 前缀
  157. "is_graded": str(is_graded_int), # 1=评级卡, 0=非评级卡
  158. "best_pt_conf": r.get("conf", ""), # best.pt 置信度
  159. "ocr_score": score, # OCR 识别的分数(无则填 authentic)
  160. "score_desc": score_desc, # 分数描述(校验用)
  161. "serial_no": serial_no if serial_no else "", # 编号(无则空)
  162. "ocr_raw": " | ".join(lines)[:200],
  163. })
  164. # 按文件名数字值排序(7位数<8位数,而不是字符串比较)
  165. def _name_sort_key(row):
  166. import re as _re
  167. nums = _re.findall(r"\d+", row["filename"])
  168. return (int(nums[0]), row["filename"]) if nums else (9999999999, row["filename"])
  169. out_rows.sort(key=_name_sort_key)
  170. # 写 CSV
  171. with open(OUT_CSV, "w", encoding="utf-8-sig", newline="") as f:
  172. w = csv.DictWriter(f, fieldnames=[
  173. "filename", "true_label", "detected_company", "is_graded", "best_pt_conf",
  174. "ocr_score", "score_desc", "serial_no", "ocr_raw"
  175. ])
  176. w.writeheader()
  177. w.writerows(out_rows)
  178. print(f"\n✓ 解析完成: {OUT_CSV} 共 {len(out_rows)} 条")
  179. # 评级卡统计
  180. n_graded = sum(1 for r in out_rows if r["is_graded"] == "1")
  181. n_ungraded = sum(1 for r in out_rows if r["is_graded"] == "0")
  182. print(f"评级卡: {n_graded}, 非评级卡: {n_ungraded}")
  183. # 公司判断准确率(只算评级卡)
  184. from collections import Counter
  185. graded_rows = [r for r in out_rows if r["is_graded"] == "1"]
  186. correct = sum(1 for r in graded_rows if r["detected_company"] == r["true_label"])
  187. print(f"\n公司判断准确率(仅评级卡 {len(graded_rows)} 张): {correct}/{len(graded_rows)} = {correct/len(graded_rows)*100:.2f}%")
  188. cnt = Counter(r["detected_company"] for r in out_rows)
  189. print("detected_company 分布(全量):")
  190. for k, v in cnt.most_common():
  191. print(f" {k or '(空=非评级卡)'}: {v}")
  192. # 分真类别的准确率(仅评级卡)
  193. from collections import defaultdict
  194. acc_by_label = defaultdict(lambda: [0, 0]) # [正确, 总]
  195. for r in graded_rows:
  196. acc_by_label[r["true_label"]][1] += 1
  197. if r["detected_company"] == r["true_label"]:
  198. acc_by_label[r["true_label"]][0] += 1
  199. print("\n分真实类别准确率(仅评级卡):")
  200. for lbl, (c, t) in sorted(acc_by_label.items()):
  201. print(f" {lbl}: {c}/{t} = {c/t*100:.2f}%")
  202. # 分数和编号统计(全量 + 评级卡)
  203. n_score = sum(1 for r in out_rows if r["ocr_score"])
  204. n_desc = sum(1 for r in out_rows if r["score_desc"])
  205. n_serial = sum(1 for r in out_rows if r["serial_no"])
  206. print(f"\n全量统计:")
  207. print(f" OCR 分数识别: {n_score}/{len(out_rows)} = {n_score/len(out_rows)*100:.1f}%")
  208. print(f" 分数描述: {n_desc}/{len(out_rows)} = {n_desc/len(out_rows)*100:.1f}%")
  209. print(f" 编号识别: {n_serial}/{len(out_rows)} = {n_serial/len(out_rows)*100:.1f}%")