# -*- 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}%")