baodan/api/insurance/ppt/golden_test.py

236 lines
8.2 KiB
Python
Raw Normal View History

"""金标 PDF 回归测试框架。
提供
- 金标 JSON 格式定义
- 提取结果与金标准的对比工具
- 字段准确率表格行召回率数字单元格准确率计算
用法
from insurance.ppt.golden_test import compare_extraction, load_gold_standard
result = compare_extraction(extracted_data, gold_standard)
print(result.summary())
"""
import json
import os
import logging
from dataclasses import dataclass, field
from typing import Optional
logger = logging.getLogger(__name__)
@dataclass
class FieldMatch:
"""单个字段的匹配结果。"""
path: str
expected: any
actual: any
match: bool # 精确匹配
near_match: bool = False # 近似匹配(数值误差 < 1%
missing: bool = False # 提取结果中缺失
@dataclass
class ComparisonResult:
"""对比结果。"""
product_name: FieldMatch = None
insured_age: FieldMatch = None
insured_gender: FieldMatch = None
currency: FieldMatch = None
annual_premium: FieldMatch = None
premium_payment_period: FieldMatch = None
sum_insured: FieldMatch = None
benefit_row_matches: list = field(default_factory=list) # list[FieldMatch]
field_accuracy: float = 0.0
row_recall: float = 0.0
number_accuracy: float = 0.0
total_fields: int = 0
matched_fields: int = 0
missing_fields: int = 0
def summary(self) -> str:
return (
f"字段准确率: {self.field_accuracy:.1%} ({self.matched_fields}/{self.total_fields})\n"
f"表格行召回率: {self.row_recall:.1%}\n"
f"数字单元格准确率: {self.number_accuracy:.1%}\n"
f"缺失字段: {self.missing_fields}"
)
def _safe_float(value) -> Optional[float]:
"""安全转为 float。"""
if value is None:
return None
try:
return float(value)
except (TypeError, ValueError):
return None
def _compare_scalar(path: str, expected, actual, tolerance: float = 0.0) -> FieldMatch:
"""对比标量字段。"""
if expected is None:
return FieldMatch(path=path, expected=expected, actual=actual, match=True)
if actual is None or actual == "" or actual == "unknown":
return FieldMatch(path=path, expected=expected, actual=actual, match=False, missing=True)
# 字符串比较
if isinstance(expected, str):
exp_lower = expected.strip().lower()
act_lower = str(actual).strip().lower()
exact = exp_lower == act_lower
near = exp_lower in act_lower or act_lower in exp_lower
return FieldMatch(path=path, expected=expected, actual=actual, match=exact, near_match=near)
# 数值比较
exp_num = _safe_float(expected)
act_num = _safe_float(actual)
if exp_num is not None and act_num is not None:
exact = abs(exp_num - act_num) < 0.01
near = abs(exp_num - act_num) / max(abs(exp_num), 1) < tolerance
return FieldMatch(path=path, expected=expected, actual=actual, match=exact, near_match=near)
return FieldMatch(path=path, expected=expected, actual=actual, match=str(expected) == str(actual))
def _compare_benefit_rows(gold_rows: list, actual_rows: list) -> tuple[list[FieldMatch], float, float]:
"""对比利益演示表行。
返回 (行匹配列表, 行召回率, 数字准确率)
"""
if not gold_rows:
return [], 1.0, 1.0
# 按 policy_year 建索引
gold_by_year = {}
for row in gold_rows:
year = row.get("policy_year") or row.get("policyYear")
if year is not None:
gold_by_year[int(year)] = row
actual_by_year = {}
for row in actual_rows:
year = row.get("policy_year") or row.get("policyYear")
if year is not None:
actual_by_year[int(year)] = row
matches = []
total_cells = 0
matched_cells = 0
found_years = 0
numeric_fields = [
"total_premium_paid", "guaranteed_cash_value", "reversionary_bonus",
"terminal_dividend", "total_surrender_value", "death_benefit",
]
for year, gold_row in gold_by_year.items():
actual_row = actual_by_year.get(year)
if actual_row is None:
matches.append(FieldMatch(
path=f"benefit_illustration[{year}]",
expected=f"{year} 年整行",
actual=None,
match=False,
missing=True,
))
continue
found_years += 1
for field_name in numeric_fields:
gold_val = gold_row.get(field_name)
if gold_val is None:
continue
total_cells += 1
actual_val = actual_row.get(field_name)
fm = _compare_scalar(
f"benefit_illustration[{year}].{field_name}",
gold_val, actual_val, tolerance=0.01,
)
if fm.match or fm.near_match:
matched_cells += 1
matches.append(fm)
row_recall = found_years / len(gold_by_year) if gold_by_year else 1.0
number_accuracy = matched_cells / total_cells if total_cells > 0 else 1.0
return matches, row_recall, number_accuracy
def compare_extraction(extracted: dict, gold: dict) -> ComparisonResult:
"""对比提取结果与金标准。
Args:
extracted: 提取结果 data 字段
gold: 金标准 JSON load_gold_standard 格式
Returns:
ComparisonResult 包含各维度准确率
"""
result = ComparisonResult()
# 标量字段对比
ext_insured = extracted.get("insured") or {}
gold_insured = gold.get("insured") or {}
ext_policy = extracted.get("policy") or {}
gold_policy = gold.get("policy") or {}
result.product_name = _compare_scalar("product_name", gold.get("product_name"), extracted.get("product_name"))
result.insured_age = _compare_scalar("insured.age", gold_insured.get("age"), ext_insured.get("age"))
result.insured_gender = _compare_scalar("insured.gender", gold_insured.get("gender"), ext_insured.get("gender"))
result.currency = _compare_scalar("policy.currency", gold_policy.get("currency"), ext_policy.get("currency"))
result.annual_premium = _compare_scalar("policy.annual_premium", gold_policy.get("annual_premium"), ext_policy.get("annual_premium"), tolerance=0.01)
result.premium_payment_period = _compare_scalar("policy.premium_payment_period", gold_policy.get("premium_payment_period"), ext_policy.get("premium_payment_period"))
result.sum_insured = _compare_scalar("policy.sum_insured", gold_policy.get("sum_insured"), ext_policy.get("sum_insured"), tolerance=0.01)
scalar_matches = [
result.product_name, result.insured_age, result.insured_gender,
result.currency, result.annual_premium, result.premium_payment_period,
result.sum_insured,
]
scalar_matches = [m for m in scalar_matches if m is not None]
# 利益表对比
gold_rows = gold.get("benefit_illustration") or []
actual_rows = extracted.get("benefit_illustration") or []
result.benefit_row_matches, result.row_recall, result.number_accuracy = _compare_benefit_rows(gold_rows, actual_rows)
# 汇总
all_matches = scalar_matches + result.benefit_row_matches
result.total_fields = len(all_matches)
result.matched_fields = sum(1 for m in all_matches if m.match or m.near_match)
result.missing_fields = sum(1 for m in all_matches if m.missing)
result.field_accuracy = result.matched_fields / result.total_fields if result.total_fields > 0 else 0.0
return result
def load_gold_standard(path: str) -> dict:
"""加载金标准 JSON 文件。
金标准格式
{
"product_name": "产品全称",
"plan_type": "savings",
"insured": {"age": 35, "gender": "male"},
"policy": {
"currency": "USD",
"annual_premium": 100000,
"premium_payment_period": 5,
"sum_insured": null
},
"benefit_illustration": [
{"policy_year": 1, "total_premium_paid": 100000, ...},
...
],
"metadata": {
"source": "公司名",
"pdf_type": "digital|scanned|mixed",
"language": "zh-CN|zh-TW|en",
"notes": "备注"
}
}
"""
with open(path, "r", encoding="utf-8") as f:
return json.load(f)