| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242 |
- # -*- coding: utf-8 -*-
- """
- parse_grading_ocr.py - 从 best.pt 检测结果 + OCR 分行文本解析出结构化字段
- 字段:
- - detected_company: 来自 best.pt label 的前缀(如 BGS-AUTHENTIC -> BGS),多框取最高置信度
- - ocr_score: OCR 识别的评分(如 10, 9.5, 8)
- - score_desc: 评分描述(可选,如 GEM MT/NM-MT/MINT,在分数上方那行)
- - serial_no: 至少 6 个连续数字字符,无则 null
- """
- import csv
- import json
- import re
- from pathlib import Path
- ROOT = Path(r"D:/顾工交接/wzj/data_all/test/grading_ocr")
- OCR_JSON = ROOT / "ocr_result.json"
- OCR_CSV = ROOT / "ocr_summary.csv"
- import os
- # 若目标文件被占用,自动换名输出
- OUT_CSV = ROOT / "grading_parsed.csv"
- if OUT_CSV.exists():
- # 尝试用 _v2/_v3 之类,直到可用
- for i in range(2, 10):
- cand = ROOT / f"grading_parsed_v{i}.csv"
- if not cand.exists():
- OUT_CSV = cand
- print(f"⚠ grading_parsed.csv 被占用,改用 {OUT_CSV.name}")
- break
- # 读 OCR JSON
- with open(OCR_JSON, "r", encoding="utf-8") as f:
- ocr_data = json.load(f)
- # 读 CSV
- rows = []
- with open(OCR_CSV, "r", encoding="utf-8-sig") as f:
- for r in csv.DictReader(f):
- rows.append(r)
- def detect_company_from_label(box_label):
- """从 best.pt label(如 'BGS-AUTHENTIC','PSA-AUTO','SGC-OLD')提取公司前缀。
- Returns:
- ("PSA", 1) 或 ("BGS", 1) 等有框
- ("", 0) 无框(非评级卡)
- """
- if not box_label or box_label == "(无检测)":
- return "", 0
- return box_label.split("-", 1)[0], 1
- # score_desc 关键词 + 对应分数(用于校验,但 score_desc 是 OCR 文本本身,不强制对应)
- DESC_TO_SCORE = [
- ("GEM MT", 10), # PSA: Gem Mint = 10
- ("GEMMINT", 10),
- ("PRISTINE", 10),
- ("GEM", 10),
- ("MINT", 9), # PSA: Mint = 9
- ("NM-MT", 8.5),
- ("NMMT", 8.5),
- ("NM", 7), # 近 Mint
- ("EX", 6),
- ("EX+", 6.5),
- ("VG-EX", 4),
- ("VG", 4),
- ("GOOD", 3),
- ("FR", 1),
- ("PR", 1.5),
- ]
- def parse_score_and_desc(lines):
- """从 OCR 分行文本里提取:
- - score: 数字分数(1-10, 或 AU)
- - score_desc: 描述(GEM MT/NM-MT/MINT 等),必须在分数上方一行
- 关键改进:
- 1. 分数限定范围 1.0-10.0(排除年份/编号)
- 2. 描述必须在分数"紧邻的上一行"(不能跨行)
- 3. 描述词必须在专门的描述关键词表里(避免"PRIZMW/C"这种卡名误匹配)
- """
- score = None
- score_desc = ""
- score_line_idx = -1
- # 找分数行: 数字必须 1-10 范围内(年份/编号/日期/超大数都排除)
- for i, ln in enumerate(lines):
- ln_clean = ln.strip()
- # 匹配 1-10 范围内的纯数字(支持 .5)
- m = re.match(r"^(\d+(\.\d+)?)$", ln_clean)
- if m:
- val = float(m.group(1))
- if 1.0 <= val <= 10.0:
- score = ln_clean
- score_line_idx = i
- break
- # 单独 "AU" 也算分数(真伪)
- if ln_clean.upper() == "AU":
- score = "AU"
- score_line_idx = i
- break
- # 描述关键词白名单(只在分数"紧邻上一行"找,且必须是这些词)
- DESC_KEYWORDS = ["GEM MT", "GEMMINT", "GEM MINT", "PRISTINE",
- "MINT", "NM-MT", "NMMT", "NM",
- "EX", "EX+", "VG-EX", "VG", "GOOD", "PR", "FR",
- "AUTHENTIC"] # AUTHENTIC 是签字的真伪描述
- # 找分数紧邻上一行的描述
- if score_line_idx > 0:
- prev_line = lines[score_line_idx - 1].strip().upper()
- # 检查上一行是否纯是描述词
- for kw in DESC_KEYWORDS:
- if kw in prev_line or prev_line == kw.replace(" ", ""):
- score_desc = lines[score_line_idx - 1].strip()
- break
- # 兜底: 如果紧邻上行没找到,扫描所有行找独立的一行(整行就是描述词)
- if not score_desc:
- for ln in lines:
- ln_stripped = ln.strip().upper()
- if ln_stripped in [k.replace(" ", "") for k in DESC_KEYWORDS]:
- score_desc = ln.strip()
- break
- return score or "", score_desc
- def parse_serial(lines):
- """找至少 6 个连续数字字符(编号)。"""
- for ln in lines:
- m = re.search(r"\d{6,}", ln)
- if m:
- return m.group(0)
- return None # 无则 null
- # best.pt 的检测结果在 infer_test.py 生成的 test_predictions.json 里
- # 我们之前下载过 test_predict/test_predictions.json,但那是 infer_test 的输出
- # 这里用 OCR_CSV 里的 filename 作为 key,通过另一份 best.pt 检测数据来获取 label
- # 实际: best.pt 的 label 在 grading_ocr.py 第一步裁剪时记录了,但 crop_records.json 里只有 box/conf,没有 label
- # 解决方案: 重新跑一遍 best.pt 拿到 label,或从测试结果里取
- # 检查有没有 best.pt label 的数据
- # 方案: 通过 test_predictions.json (infer_test 生成的) 拿到 detected label
- PRED_JSON = Path(r"D:/顾工交接/wzj/data_all/test/test_predict/test_predictions.json")
- pred_data = {}
- if PRED_JSON.exists():
- with open(PRED_JSON, "r", encoding="utf-8") as f:
- pred_data = json.load(f)
- # pred_data 格式: {"records": [{true, pred, conf, file}, ...]}
- records = pred_data.get("records", [])
- pred_label = {r["file"]: r["pred"] for r in records}
- print(f"✓ 从 test_predictions.json 加载 best.pt 预测结果: {len(pred_label)} 条")
- else:
- pred_label = {}
- print("⚠ test_predictions.json 不存在,detected_company 字段将为空")
- # 解析
- out_rows = []
- for r in rows:
- name = r["filename"]
- lines = ocr_data.get(name, {}).get("lines", [])
- box_label = pred_label.get(name, "")
- detected_company, is_graded_int = detect_company_from_label(box_label)
- score, score_desc = parse_score_and_desc(lines)
- serial_no = parse_serial(lines)
- # 评级卡但 OCR 没识别出分数 -> 填 "authentic"(说明卡是真的评分过,但具体分数OCR漏识)
- if is_graded_int == 1 and not score:
- score = "authentic"
- # 非评级卡 -> 一律清空(OCR在原图上抓到的数字是噪音,不是评级分数)
- if is_graded_int == 0:
- score = ""
- out_rows.append({
- "filename": name,
- "true_label": r["true_label"],
- "detected_company": detected_company, # 来自 best.pt label 前缀
- "is_graded": str(is_graded_int), # 1=评级卡, 0=非评级卡
- "best_pt_conf": r.get("conf", ""), # best.pt 置信度
- "ocr_score": score, # OCR 识别的分数(无则填 authentic)
- "score_desc": score_desc, # 分数描述(校验用)
- "serial_no": serial_no if serial_no else "", # 编号(无则空)
- "ocr_raw": " | ".join(lines)[:200],
- })
- # 按文件名数字值排序(7位数<8位数,而不是字符串比较)
- def _name_sort_key(row):
- import re as _re
- nums = _re.findall(r"\d+", row["filename"])
- return (int(nums[0]), row["filename"]) if nums else (9999999999, row["filename"])
- out_rows.sort(key=_name_sort_key)
- # 写 CSV
- with open(OUT_CSV, "w", encoding="utf-8-sig", newline="") as f:
- w = csv.DictWriter(f, fieldnames=[
- "filename", "true_label", "detected_company", "is_graded", "best_pt_conf",
- "ocr_score", "score_desc", "serial_no", "ocr_raw"
- ])
- w.writeheader()
- w.writerows(out_rows)
- print(f"\n✓ 解析完成: {OUT_CSV} 共 {len(out_rows)} 条")
- # 评级卡统计
- n_graded = sum(1 for r in out_rows if r["is_graded"] == "1")
- n_ungraded = sum(1 for r in out_rows if r["is_graded"] == "0")
- print(f"评级卡: {n_graded}, 非评级卡: {n_ungraded}")
- # 公司判断准确率(只算评级卡)
- from collections import Counter
- graded_rows = [r for r in out_rows if r["is_graded"] == "1"]
- correct = sum(1 for r in graded_rows if r["detected_company"] == r["true_label"])
- print(f"\n公司判断准确率(仅评级卡 {len(graded_rows)} 张): {correct}/{len(graded_rows)} = {correct/len(graded_rows)*100:.2f}%")
- cnt = Counter(r["detected_company"] for r in out_rows)
- print("detected_company 分布(全量):")
- for k, v in cnt.most_common():
- print(f" {k or '(空=非评级卡)'}: {v}")
- # 分真类别的准确率(仅评级卡)
- from collections import defaultdict
- acc_by_label = defaultdict(lambda: [0, 0]) # [正确, 总]
- for r in graded_rows:
- acc_by_label[r["true_label"]][1] += 1
- if r["detected_company"] == r["true_label"]:
- acc_by_label[r["true_label"]][0] += 1
- print("\n分真实类别准确率(仅评级卡):")
- for lbl, (c, t) in sorted(acc_by_label.items()):
- print(f" {lbl}: {c}/{t} = {c/t*100:.2f}%")
- # 分数和编号统计(全量 + 评级卡)
- n_score = sum(1 for r in out_rows if r["ocr_score"])
- n_desc = sum(1 for r in out_rows if r["score_desc"])
- n_serial = sum(1 for r in out_rows if r["serial_no"])
- print(f"\n全量统计:")
- print(f" OCR 分数识别: {n_score}/{len(out_rows)} = {n_score/len(out_rows)*100:.1f}%")
- print(f" 分数描述: {n_desc}/{len(out_rows)} = {n_desc/len(out_rows)*100:.1f}%")
- print(f" 编号识别: {n_serial}/{len(out_rows)} = {n_serial/len(out_rows)*100:.1f}%")
|