aiflow-math-ink-06-intermediate / scripts /audit_math_ink_06_federation_errors.py
cwLeeDev's picture
Publish case-context proxy rejection and source error audit
4dee47b verified
Raw
History Blame Contribute Delete
10.7 kB
"""교정된 Math Ink 0.6 federation checkpoint의 source·writer·label 병목을 감사한다."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import sys
import torch
from torch.utils.data import DataLoader
PROJECT_ROOT = Path(__file__).parents[1]
SOURCE_ROOT = PROJECT_ROOT / "src"
if str(SOURCE_ROOT) not in sys.path:
sys.path.insert(0, str(SOURCE_ROOT))
if str(PROJECT_ROOT / "scripts") not in sys.path:
sys.path.insert(0, str(PROJECT_ROOT / "scripts"))
from audit_math_ink_06_online_errors import summarize_online_predictions06
from math_grid_drawer.research.case_context06 import (
SIZE_DEPENDENT_CASE_BASES06,
apply_relative_case_context06,
)
from math_grid_drawer.research.ink06_federation import (
FederatedPairedInk06Dataset,
load_product_federation06,
resolve_training_device06,
)
from math_grid_drawer.research.math_ink_06 import MathInk06Engine
def _bbox_context06(record: dict) -> dict:
"""필요 변수: 원본 source stroke·canvas. 작동 원리: 128 정규화 전에 있던 상대 높이를 context resolver용 bbox로 보존한다."""
points = []
for stroke in record.get("strokes") or []:
values = stroke.get("points", []) if isinstance(stroke, dict) else stroke
for point in values:
if isinstance(point, dict):
points.append((float(point["x"]), float(point["y"])))
else:
points.append((float(point[0]), float(point[1])))
if not points:
return dict(record)
canvas = record.get("canvas") or {}
height = max(float(canvas.get("height", 1.0)), 1e-6)
top = min(point[1] for point in points) / height
bottom = max(point[1] for point in points) / height
output = dict(record)
output["baseline_context"] = {
"bbox_top": top,
"bbox_bottom": bottom,
"bbox_height": max(bottom - top, 1e-6),
}
return output
def _case_size_proxy06(
*, logits: torch.Tensor, labels: tuple[str, ...], records: list[dict],
) -> dict:
"""필요 변수: writer별 logits·원본 bbox. 작동 원리: 동일 writer를 행 anchor의 대리값으로 삼아 대소문자 상대크기 이득과 회귀를 함께 센다."""
enriched = [_bbox_context06(record) for record in records]
groups: dict[tuple[str, str], list[int]] = {}
for index, record in enumerate(enriched):
key = (str(record.get("source")), str(record.get("writer_key") or "missing"))
groups.setdefault(key, []).append(index)
predictions = logits.argmax(dim=1)
source_rows: dict[str, dict[str, int]] = {}
for (source, _writer), indices in groups.items():
group_logits = logits[indices]
group_records = [enriched[index] for index in indices]
resolved, decisions = apply_relative_case_context06(group_logits, labels, group_records)
raw = group_logits.argmax(dim=1)
revised = resolved.argmax(dim=1)
for local_index, global_index in enumerate(indices):
truth = str(records[global_index]["label"])
if len(truth) != 1 or not truth.isascii() or truth.lower() not in SIZE_DEPENDENT_CASE_BASES06:
continue
row = source_rows.setdefault(source, {
"samples": 0, "raw_correct": 0, "resolved_correct": 0,
"changed": 0, "beneficial": 0, "harmful": 0,
})
raw_label = labels[int(raw[local_index])]
resolved_label = labels[int(revised[local_index])]
row["samples"] += 1
row["raw_correct"] += int(raw_label == truth)
row["resolved_correct"] += int(resolved_label == truth)
row["changed"] += int(raw_label != resolved_label)
row["beneficial"] += int(raw_label != truth and resolved_label == truth)
row["harmful"] += int(raw_label == truth and resolved_label != truth)
by_source = {
source: {
**row,
"raw_top1": row["raw_correct"] / row["samples"] if row["samples"] else 0.0,
"resolved_top1": row["resolved_correct"] / row["samples"] if row["samples"] else 0.0,
"delta_percentage_points": (
(row["resolved_correct"] - row["raw_correct"]) * 100.0 / row["samples"]
if row["samples"] else 0.0
),
}
for source, row in sorted(source_rows.items())
}
return {
"scope": "same-writer isolated-glyph relative-size proxy; not continuous-formula product evidence",
"writers": len(groups),
"by_source": by_source,
"product_validation": False,
}
def _evaluate_online06(
engine: MathInk06Engine,
records: list[dict],
*,
batch_size: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""필요 변수: 교정 test record·engine. 작동 원리: 동일 분모의 정답과 exact log-probability를 CPU에 모은다."""
exact_to_index = {label: index for index, label in enumerate(engine.labels)}
family_to_index = {label: index for index, label in enumerate(engine.family_labels)}
loader = DataLoader(
FederatedPairedInk06Dataset(records, exact_to_index, family_to_index),
batch_size=batch_size,
shuffle=False,
num_workers=0,
)
targets: list[torch.Tensor] = []
probabilities: list[torch.Tensor] = []
engine.model.eval()
with torch.inference_mode():
for online, _raster, _coordinates, _states, target, _family, _source in loader:
logits, family_logits = engine.model.forward_online(online.to(engine.device))
fused = engine._fuse_online_exact06(logits, family_logits)
targets.append(target.cpu())
probabilities.append(fused.log_softmax(dim=1).cpu())
return torch.cat(targets), torch.cat(probabilities)
def main() -> None:
"""필요 변수: provenance checkpoint·승인 federation. 작동 원리: 실제 명시 test의 source/writer/label 오류를 JSON으로 고정한다."""
parser = argparse.ArgumentParser(description="Audit clean federation exact bottlenecks")
parser.add_argument("--checkpoint", type=Path, required=True)
parser.add_argument("--registry", type=Path, default=PROJECT_ROOT / "research/dataset_registry.json")
parser.add_argument("--source-registry", type=Path, default=PROJECT_ROOT / "research/math_ink_06_source_registry.json")
parser.add_argument("--commercial", type=Path, default=PROJECT_ROOT / "research/data/external_trajectory_v1/commercial_ccby4.jsonl.gz")
parser.add_argument("--hwrt", type=Path, default=PROJECT_ROOT / "research/data/open_pretrain/hwrt_expanded_v2/hwrt_expanded.jsonl.gz")
parser.add_argument("--approval", type=Path, default=PROJECT_ROOT / "research/approvals/HWRT-ODBL-USE-APPROVAL-v1.json")
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--batch-size", type=int, default=32)
parser.add_argument("--maximum-per-source", type=int, default=0)
parser.add_argument("--device", default="auto")
args = parser.parse_args()
device = resolve_training_device06(args.device)
engine = MathInk06Engine(args.checkpoint, device=device)
sources = load_product_federation06(
registry_path=args.registry,
commercial_path=args.commercial,
hwrt_path=args.hwrt,
approval_path=args.approval,
allowed_labels=engine.labels,
source_registry_path=args.source_registry,
)
training_metadata = [
record for source in sources for record in source.records if record.get("eligible_for_training")
]
test_records: list[dict] = []
for source in sources:
rows = [record for record in source.records if record.get("split") == "test"]
if args.maximum_per_source > 0:
rows = rows[:args.maximum_per_source]
test_records.extend(rows)
targets, log_probabilities = _evaluate_online06(
engine,
test_records,
batch_size=args.batch_size,
)
metadata = [{
**record,
"writer_key": str(record.get("writer_key") or record.get("writer_id") or "missing"),
} for record in test_records]
report = summarize_online_predictions06(
labels=engine.labels,
metadata=metadata,
training_metadata=training_metadata,
targets=targets,
log_probabilities=log_probabilities,
exact_family_index=engine.exact_family_index.cpu(),
)
source_bottlenecks = {}
for source_id in sorted({str(record["source"]) for record in metadata}):
indices = [
index for index, record in enumerate(metadata)
if str(record["source"]) == source_id
]
index_tensor = torch.tensor(indices, dtype=torch.long)
source_report = summarize_online_predictions06(
labels=engine.labels,
metadata=[metadata[index] for index in indices],
training_metadata=[
record for record in training_metadata
if str(record["source"]) == source_id
],
targets=targets.index_select(0, index_tensor),
log_probabilities=log_probabilities.index_select(0, index_tensor),
exact_family_index=engine.exact_family_index.cpu(),
)
source_bottlenecks[source_id] = {
key: source_report[key]
for key in (
"samples", "top1", "top5", "shape_family_top1",
"same_family_error_rate", "writer_accuracy_p10",
"writer_accuracy_minimum", "top_confusions",
"highest_error_labels", "lowest_supported_labels",
)
}
report["source_bottlenecks"] = source_bottlenecks
report["case_size_proxy"] = _case_size_proxy06(
logits=log_probabilities,
labels=engine.labels,
records=metadata,
)
report.update({
"checkpoint": str(args.checkpoint),
"device": device,
"training_source_ids": list(torch.load(
args.checkpoint,
map_location="cpu",
weights_only=False,
).get("training_source_ids", [])),
"product_validation": False,
})
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(
json.dumps(report, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
print(json.dumps({
key: report[key]
for key in (
"samples", "top1", "top5", "shape_family_top1",
"same_family_error_rate", "source_metrics", "writer_accuracy_p10",
"writer_accuracy_minimum", "top_confusions",
)
}, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()