baodan/api/insurance/ppt/golden_test.py

236 lines
8.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""金标 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)