diff --git a/.github/workflows/benchmark.yml b/.github/workflows/benchmark.yml index 924663d..68c13e9 100644 --- a/.github/workflows/benchmark.yml +++ b/.github/workflows/benchmark.yml @@ -10,6 +10,7 @@ on: paths: - "src/**" - "benchmarks/measure_recall.py" + - "benchmarks/benchmark.py" - "scripts/**" - "pyproject.toml" - ".github/workflows/benchmark.yml" diff --git a/benchmarks/benchmark.py b/benchmarks/benchmark.py new file mode 100644 index 0000000..a7ef83b --- /dev/null +++ b/benchmarks/benchmark.py @@ -0,0 +1,795 @@ +""" +benchmark.py — baseline(Docling HybridChunker) vs EvidenceChunker 벤치마크 + +유일한 평가 엔진. measure_recall.py는 이 파일의 run()을 호출하는 CI 전용 +얇은 래퍼(DagsHub/MLflow 로깅만 담당)이므로, 채점 로직은 여기만 고치면 된다. + +지표: PageHit@1/5/10(모든 arm 동일 기준), evidence_hit@1(정답 값+문맥 키 +포함 여부), TableHit@k(정답 표 자체 일치, 페이지 일치보다 엄격), ablation +(caption/context 제거), 문서 단위 부트스트랩 CI + McNemar, 동일 토큰 예산 +LLM 실답변 채점(Gemini). + +사전 준비(직접 설치할 것 — 이 스크립트는 설치를 수행하지 않는다): + pip install docling docling-core sentence-transformers langchain-core \\ + google-genai tiktoken + pip install -e . + +사용법: + python benchmark.py --pdf-dir ./data/pdfs --qa-dir ./auto_qa \\ + --out-dir ./results --dev-only + python benchmark.py --pdf-dir ./data/pdfs --qa-dir ./auto_qa \\ + --out-dir ./results --gemini-api-key $GEMINI_API_KEY --llm-sample-n 300 + + GEMINI_API_KEY가 없으면 LLM 실답변 평가만 자동으로 건너뛴다. + CI에서는 measure_recall.py를 통해 실행된다. +""" + +from __future__ import annotations + +import argparse +import copy +import gc +import json +import math +import os +import re +import time +from collections import defaultdict +from pathlib import Path +from typing import Optional + +import numpy as np + +EMBED_MODEL_NAME = "sentence-transformers/all-MiniLM-L6-v2" +ENCODE_BATCH = 256 +BBOX_THRESHOLD = 300.0 +SIM_THRESHOLD = 0.00 +K_LIST = (1, 5, 10) +TOP_K_MAX = max(K_LIST) + +CONTEXT_TOKEN_BUDGET = 1500 +LLM_SAMPLE_N = 300 +LLM_SAMPLE_SEED = 20260915 +MODEL_CANDIDATES = [ + "gemini-3.5-flash-lite", + "gemini-3.5-flash", + "gemini-flash-latest", + "gemini-2.5-flash", +] + +# dev 서브셋 — v04와 동일 문서 구성 +DEV_DOCS = [ + "1. Attention is all you need", "6. DPO", "14. CLIP", "21. T5", "26. MMLU", + "35. APB", "39. risk sharing", "42. rural housing", "45. gao-25-107649", + "47. gao-26-107681", "50. gao-26-107884", "52. gao-26-108011", + "55. gao-26-108116", "60. ieee1", "64. ieee5", "66. ieee7", "70. ieee11", + "72. ieee13", "80. ssrn-1331573", "85. ssrn-2760631", +] +_LEADING_NUM = re.compile(r"^(\d+)\.") + + +# =========================================================================== +# 1. 채점 함수 +# =========================================================================== + +def normalize_for_em(s: str) -> str: + if not s: + return "" + s = s.lower().replace("\xa0", " ") + s = re.sub(r"[^\w.\-±× ]+", " ", s) + s = re.sub(r"\s+", " ", s).strip() + s = re.sub(r"\.$", "", s) + return s + + +def has_token(needle: str, haystack_norm: str) -> bool: + n = normalize_for_em(needle) + if not n: + return False + return re.search(r"(? bool: + """정답 값 + 모든 context_keys가 같은 청크 안에 있으면 True. 구 이름 em_hit.""" + if not spec or not chunk_text: + return False + c = normalize_for_em(chunk_text) + if not has_token(spec.get("value", ""), c): + return False + return all(has_token(k, c) for k in spec.get("context_keys", [])) + + +def classify_real_driver(b_ok: bool, e_ok: bool) -> str: + if b_ok and e_ok: + return "both_right" + if b_ok and not e_ok: + return "baseline_win_eu_lose" + if not b_ok and e_ok: + return "eu_win_baseline_lose" + return "both_wrong" + + +def paired_statistics(doc_ids, baseline, treatment, repeats: int = 50000, seed: int = 20260915) -> dict: + """문서 단위 클러스터 부트스트랩 CI(주 지표) + McNemar 검정(보조 지표).""" + b = np.asarray(list(baseline), dtype=np.int64) + e = np.asarray(list(treatment), dtype=np.int64) + doc_ids = list(doc_ids) + if len(doc_ids) != len(b) or len(b) != len(e) or not len(b): + raise ValueError("Paired vectors must be nonempty and equal length.") + + grouped = defaultdict(lambda: [0, 0]) + for d, delta in zip(doc_ids, e - b): + grouped[d][0] += 1 + grouped[d][1] += int(delta) + a = np.asarray([grouped[d] for d in sorted(grouped)], dtype=np.int64) + + ci = None + if len(a) >= 2: + rng, values = np.random.default_rng(seed), [] + for start in range(0, repeats, 2048): + idx = rng.integers(0, len(a), size=(min(2048, repeats - start), len(a))) + total = a[idx].sum(axis=1) + values.append(100 * total[:, 1] / total[:, 0]) + ci = np.quantile(np.concatenate(values), [.025, .975]).tolist() + + b_only = int(((b == 1) & (e == 0)).sum()) + e_only = int(((b == 0) & (e == 1)).sum()) + discordant = b_only + e_only + chi2 = max(abs(b_only - e_only) - 1, 0) ** 2 / discordant if discordant else 0.0 + p = math.erfc(math.sqrt(chi2 / 2)) if discordant else 1.0 + + return { + "n_questions": len(b), "n_documents": len(a), + "baseline": float(b.mean()), "treatment": float(e.mean()), + "difference_pp": float(100 * (e - b).mean()), + "cluster_bootstrap_ci95_pp": ci, + "bootstrap": {"unit": "document", "statistic": "micro_rate_difference", + "method": "percentile", "repeats": repeats, "seed": seed}, + "mcnemar_supplementary": { + "method": "chi_square_continuity_corrected", "chi2": chi2, "p_value": p, + "baseline_only": b_only, "treatment_only": e_only, + "assumption": "질문 간 독립 가정 — 문서 내 의존성은 보정하지 않음(보조 지표)", + }, + "ci_note": ("문서를 독립 표집 단위로 취급한 근사치이며, LLM 생성 답변 정확도의 CI가 아님" + if ci is not None else "문서 2개 미만이라 CI 계산 불가"), + } + + +def _self_test() -> None: + """Docling/GPU 없이 채점 로직만 검증. import 시 자동 실행.""" + spec = {"value": "42.3", "context_keys": ["APAC"]} + assert evidence_hit(spec, "APAC region | Q2: 42.3") is True + assert evidence_hit(spec, "42.3 percent, unrelated") is False + assert evidence_hit(spec, "APAC region only, no numbers") is False + + assert classify_real_driver(True, True) == "both_right" + assert classify_real_driver(True, False) == "baseline_win_eu_lose" + assert classify_real_driver(False, True) == "eu_win_baseline_lose" + assert classify_real_driver(False, False) == "both_wrong" + + doc_ids, b, e = [], [], [] + for _ in range(1181): doc_ids.append("d0"); b.append(1); e.append(1) + for _ in range(521): doc_ids.append("d0"); b.append(0); e.append(1) + for _ in range(382): doc_ids.append("d0"); b.append(1); e.append(0) + for _ in range(641): doc_ids.append("d0"); b.append(0); e.append(0) + r = paired_statistics(doc_ids, b, e, repeats=2000) + assert abs(r["mcnemar_supplementary"]["chi2"] - 21.0897) < 0.01, r + assert abs(r["mcnemar_supplementary"]["p_value"] - 4.3828e-06) < 1e-8, r + assert r["cluster_bootstrap_ci95_pp"] is None + + +_self_test() + + +# =========================================================================== +# 2. PDF <-> QA 매칭 +# =========================================================================== + +def _sort_key(p: Path): + m = _LEADING_NUM.match(p.name) + return (int(m.group(1)) if m else 10 ** 6, p.name) + + +def pdf_qa_pairs(pdf_dir: Path, qa_dir: Path, dev_only: bool, max_pdfs): + pairs, unmatched, empty = [], [], [] + for qa_path in sorted(qa_dir.glob("*_qa.json"), key=_sort_key): + if qa_path.name.startswith("_"): + continue + doc_id = qa_path.stem[:-3] if qa_path.stem.endswith("_qa") else qa_path.stem + if dev_only and doc_id not in DEV_DOCS: + continue + pdf_path = pdf_dir / f"{doc_id}.pdf" + if not pdf_path.exists(): + unmatched.append(qa_path.name) + continue + try: + n = len(json.load(open(qa_path, encoding="utf-8"))) + except Exception: + n = 0 + if n == 0: + empty.append(qa_path.name) + continue + pairs.append((pdf_path, qa_path, doc_id)) + + if unmatched: + print(f" [warn] PDF 매칭 실패 {len(unmatched)}건: {unmatched[:5]}") + if empty: + print(f" [skip] 0문항 파일 {len(empty)}건: {empty}") + print(f" [pairs] {len(pairs)} PDF+QA") + return pairs[:max_pdfs] if max_pdfs else pairs + + +# =========================================================================== +# 3. 코퍼스 구성 — base(hybrid)/EU + ablation 3종 +# =========================================================================== + +_converter = None + + +def get_converter(): + global _converter + if _converter is None: + from docling.document_converter import DocumentConverter, PdfFormatOption + from docling.datamodel.pipeline_options import PdfPipelineOptions + # 라이브러리의 make_converter()와 동일하게 명시 고정 — Docling 기본값에 + # 기대면 버전이 바뀔 때 표 구조 추출이 조용히 꺼질 수 있다. + pipeline_options = PdfPipelineOptions(do_table_structure=True) + _converter = DocumentConverter( + format_options={"pdf": PdfFormatOption(pipeline_options=pipeline_options)} + ) + return _converter + + +def chunk_page(c): + try: + for di in c.meta.doc_items: + if di.prov: + return di.prov[0].page_no + except Exception: + pass + return None + + +def chunk_pages_all(c): + """청크가 걸친 모든 페이지 집합 (PageHit@k 판정용).""" + pages = set() + try: + for di in c.meta.doc_items: + for p in (di.prov or []): + pages.add(p.page_no) + except Exception: + pass + return pages + + +def make_ablation_units(eu_list, drop_caption: bool, drop_context: bool): + out = [] + for eu in eu_list: + eu2 = copy.deepcopy(eu) + if drop_caption: + eu2.caption_text = None + if drop_context: + eu2.context_before = [] + eu2.context_after = [] + out.append(eu2) + return out + + +def build_all_corpora(pdf_path: Path, doc_id: str): + """PDF 1회 파싱으로 baseline + EU + ablation arm 전부 생성.""" + from docling.chunking import HybridChunker + from docling_core.types.doc import DocItemLabel + from evidence_chunker.chunker import build_evidence_units + from evidence_chunker.split import split_oversized_units + from evidence_chunker.parser.docling import DoclingParser + from evidence_chunker.export import TextChunk, filter_consumed_paragraphs + + doc = get_converter().convert(str(pdf_path)).document + + all_chunks = list(HybridChunker().chunk(doc)) + is_table_chunk = lambda c: any(di.label == DocItemLabel.TABLE for di in c.meta.doc_items) + + parsed = DoclingParser().from_doc(doc) + eu_list = build_evidence_units(parsed, BBOX_THRESHOLD, SIM_THRESHOLD, doc_id) + eu_list = split_oversized_units(eu_list) + + # base 팔의 표 청크에도 table_index를 매겨 TableHit@k를 양쪽 동일 기준으로 + # 판정한다(페이지 일치만으론 같은 페이지의 다른 표와 구분 불가). doc.tables[i]가 + # parsed.tables[i]와 동일 순서/객체라 identity로 매칭. + base_table_index = [None] * len(all_chunks) + for ci, c in enumerate(all_chunks): + if not is_table_chunk(c): + continue + for ti, t in enumerate(doc.tables): + if any(di is t for di in c.meta.doc_items): + base_table_index[ci] = ti + break + + non_table = [c for c in all_chunks if not is_table_chunk(c)] + non_table_kept = filter_consumed_paragraphs(non_table, eu_list) + + def _pack(eu_variant): + # TextChunk는 page_no 1개만 저장해 base(chunk_pages_all=전체 페이지)와 + # PageHit 기준이 비대칭해진다 — page_span을 명시로 채워 맞춘다. + packed = [] + for i, c in enumerate(non_table_kept): + tc = TextChunk(c, doc_id, i) + tc.metadata["page_span"] = chunk_pages_all(c) + packed.append(tc) + return eu_variant + packed + + arms = { + "base": {"chunks": all_chunks, "kind": "hybrid", "table_idx": base_table_index}, + "eu": {"chunks": _pack(eu_list), "kind": "eu"}, + "eu_no_caption": {"chunks": _pack(make_ablation_units(eu_list, True, False)), "kind": "eu"}, + "eu_no_context": {"chunks": _pack(make_ablation_units(eu_list, False, True)), "kind": "eu"}, + "eu_no_both": {"chunks": _pack(make_ablation_units(eu_list, True, True)), "kind": "eu"}, + } + stats = { + "n_tables": len(parsed.tables), "n_eu": len(eu_list), + "n_split": sum(1 for eu in eu_list if eu.is_split), + "n_baseline_chunks": len(all_chunks), + } + del doc + gc.collect() + return arms, stats + + +# =========================================================================== +# 4. 랭킹 및 PageHit@k / TableHit@k — 모든 arm에 동일 기준 적용 +# =========================================================================== + +def rank_indices(scores, k): + return np.argsort(-np.asarray(scores), kind="stable")[:k].tolist() + + +class _MPDoc: + """dedupe_by_chunk_id()용 최소 어댑터(.metadata만 필요).""" + def __init__(self, idx, chunk_id): + self.idx = idx + self.metadata = {"chunk_id": chunk_id} + + +def top_k_with_maxpool(scores, chunk_ids, k, is_eu_kind): + """EU 계열이면 max-pool dedup 적용, baseline이면 그냥 top-k.""" + wide = rank_indices(scores, max(k * 4, 20)) + if not is_eu_kind: + return wide[:k] + from evidence_chunker.export.langchain import dedupe_by_chunk_id + candidates = [(_MPDoc(i, chunk_ids[i]), float(scores[i])) for i in wide] + selected = dedupe_by_chunk_id(candidates, k=k) + return [d.idx for d, _ in selected] + + +def page_hit_at_k(order, page_sets, gold_page): + return {k: any(gold_page in page_sets[i] for i in order[:k]) for k in K_LIST} + + +def table_hit_at_k(order, table_idx_list, gold_table_index): + """정답 표(gold_table_index) 자체를 top-k에서 찾았는가. gold가 None(표 없는 질문)이면 호출하지 않는다.""" + return {k: any(table_idx_list[i] == gold_table_index for i in order[:k]) for k in K_LIST} + + +# =========================================================================== +# 5. 동일 토큰 예산 + LLM 실답변 평가 +# =========================================================================== + +_gemini_client = None +_working_model = {"name": None} + + +def get_gemini_client(api_key: str): + global _gemini_client + if _gemini_client is None: + from google import genai + _gemini_client = genai.Client(api_key=api_key, + http_options=genai.types.HttpOptions(timeout=30_000)) + return _gemini_client + + +def ask_llm(context: str, question: str, api_key: Optional[str]) -> str: + if not api_key: + return "" + client = get_gemini_client(api_key) + prompt = ("Answer the question using only the context below. " + "If the context does not contain the answer, say so explicitly.\n\n" + f"Context:\n{context}\n\nQuestion: {question}") + cached = _working_model["name"] + candidates = ([cached] if cached else []) + [m for m in MODEL_CANDIDATES if m != cached] + RETRYABLE = ("503", "504", "429") + attempts = [] + for model_name in candidates[:2]: + for backoff in (0, 2, 4): + if backoff: + time.sleep(backoff) + try: + resp = client.models.generate_content(model=model_name, contents=prompt) + _working_model["name"] = model_name + return resp.text.strip() + except Exception as e: + err = str(e) + attempts.append((model_name, err.split(chr(10))[0][:100])) + if not any(code in err for code in RETRYABLE): + break + return f"[LLM 호출 실패] {attempts[-1] if attempts else 'unknown'}" + + +def build_budget_context(order, texts, budget_tokens: int, tokenizer=None) -> str: + def _count(s): + return len(tokenizer.encode(s)) if tokenizer else len(s.split()) + + parts, used = [], 0 + for i in order: + t = texts[i] + n = _count(t) + if used + n > budget_tokens and parts: + break + parts.append(t) + used += n + return "\n\n---\n\n".join(parts) + + +def stratified_sample(rows, n, seed=LLM_SAMPLE_SEED): + rng = np.random.default_rng(seed) + by_type = defaultdict(list) + for r in rows: + by_type[r["type"]].append(r) + total = len(rows) + picked = [] + for t, group in by_type.items(): + take = max(1, round(n * len(group) / total)) if total else 0 + idx = rng.choice(len(group), size=min(take, len(group)), replace=False) + picked += [group[i] for i in idx] + return picked[:n] + + +def run_llm_eval(rows, all_arm_data, api_key: Optional[str], llm_sample_n: int, tokenizer=None): + if not api_key: + print("[LLM 평가] GEMINI_API_KEY 없음 — 건너뜀") + return [] + + sample = stratified_sample(rows, llm_sample_n) + print(f"[LLM 평가] {len(sample)}문항 샘플로 baseline/EU 각각 채점") + + llm_rows = [] + for j, r in enumerate(sample, 1): + ad = all_arm_data.get(r["doc_id"]) + if not ad or "base" not in ad or "eu" not in ad: + continue + b_order = r["arms"]["base"]["order"] + e_order = r["arms"]["eu"]["order"] + b_ctx = build_budget_context(b_order, ad["base"]["texts"], CONTEXT_TOKEN_BUDGET, tokenizer) + e_disp = ad.get("eu__disp", ad["eu"]["texts"]) + e_ctx = build_budget_context(e_order, e_disp, CONTEXT_TOKEN_BUDGET, tokenizer) + + b_answer = ask_llm(b_ctx, r["question"], api_key) + time.sleep(2) # 무료 티어 분당 제한 대비 + e_answer = ask_llm(e_ctx, r["question"], api_key) + time.sleep(2) + + llm_rows.append({ + "doc_id": r["doc_id"], "qid": r["qid"], "type": r["type"], + "baseline_llm_correct": evidence_hit(r["answer_spec"], b_answer), + "eu_llm_correct": evidence_hit(r["answer_spec"], e_answer), + "baseline_answer": b_answer, "eu_answer": e_answer, + }) + if j % 25 == 0: + print(f" [{j}/{len(sample)}]") + return llm_rows + + +# =========================================================================== +# 6. 문서 1개 평가 +# =========================================================================== + +def evaluate_document(pdf_path: Path, qa_path: Path, doc_id: str, model, tokenizer=None): + print(f"\n{'-'*62}\n {doc_id}") + try: + arms, stats = build_all_corpora(pdf_path, doc_id) + except Exception as e: + print(f" [ERR] {type(e).__name__}: {e}") + return [] + + print(f" 표 {stats['n_tables']} -> EU {stats['n_eu']} (분할 {stats['n_split']}) " + f"baseline 청크 {stats['n_baseline_chunks']}") + + arm_data = {} + for name, a in arms.items(): + chunks = a["chunks"] + if not chunks: + continue + if a["kind"] == "hybrid": + texts = [c.text for c in chunks] + pages = [chunk_pages_all(c) for c in chunks] + chunk_ids = [f"{doc_id}-hybrid-{i}" for i in range(len(chunks))] + table_idx = a.get("table_idx") or [None] * len(chunks) + else: + # EU 계열: retrieval_units(행 단위)로 색인, 표시는 EU 전체 text + units, ids, pages_u, table_idx_u, disp = [], [], [], [], [] + for c in chunks: + page = c.metadata.get("page_no") + page_span = c.metadata.get("page_span") or ({page} if page is not None else set()) + c_table_idx = c.metadata.get("table_index") + for u in c.retrieval_units: + units.append(u) + ids.append(c.chunk_id) + pages_u.append(set(page_span) if page_span else {page}) + table_idx_u.append(c_table_idx) + disp.append(c.text) + texts, pages, chunk_ids, table_idx = units, pages_u, ids, table_idx_u + arm_data[name + "__disp"] = disp + if not texts: + continue + emb = model.encode(texts, normalize_embeddings=True, show_progress_bar=False, batch_size=ENCODE_BATCH) + arm_data[name] = {"texts": texts, "pages": pages, "chunk_ids": chunk_ids, + "table_idx": table_idx, "emb": emb, "kind": a["kind"]} + + # 표 0개 회귀 검사(코퍼스 레벨): eu_list가 비면 eu 팔 텍스트/페이지는 base와 + # 완전히 같아야 한다. 불일치는 리트리버 차이가 아니라 채점 입력 자체가 + # 어긋났다는 신호. + if stats["n_tables"] == 0 and "base" in arm_data and "eu" in arm_data: + b_texts, e_texts = arm_data["base"]["texts"], arm_data["eu"]["texts"] + b_pages, e_pages = arm_data["base"]["pages"], arm_data["eu"]["pages"] + if b_texts != e_texts: + print(f" [회귀경고] {doc_id}: 표 0개인데 base/eu 텍스트 리스트 자체가 다름 " + f"(base={len(b_texts)}개 eu={len(e_texts)}개)") + elif b_pages != e_pages: + mismatches = [i for i in range(len(b_pages)) if b_pages[i] != e_pages[i]] + print(f" [회귀경고] {doc_id}: 표 0개인데 페이지 집합이 {len(mismatches)}개 청크에서 불일치 " + f"(예: idx={mismatches[0]} base={b_pages[mismatches[0]]} eu={e_pages[mismatches[0]]})") + + qa_list = json.load(open(qa_path, encoding="utf-8")) + usable = [q for q in qa_list if q.get("question_delabeled") is not None + and q.get("subset", "main") in (None, "main")] + if not usable or "base" not in arm_data or "eu" not in arm_data: + print(" [warn] 사용 가능 문항 0 또는 코퍼스 비어있음") + return [] + + q_emb = model.encode([q["question_delabeled"] for q in usable], + normalize_embeddings=True, show_progress_bar=False, batch_size=ENCODE_BATCH) + + rows = [] + for i, qa in enumerate(usable): + gold_page = qa.get("page") + row = {"doc_id": doc_id, "qid": qa.get("qid", ""), "type": qa.get("type", "unknown"), + "question": qa["question_delabeled"], "expected_page": gold_page, + "answer_spec": qa.get("answer_spec"), "answer": qa.get("answer", ""), "arms": {}, + "meta": qa.get("meta") or {}} + + gold_table_index = row["meta"].get("table_index") + row["gold_table_index"] = gold_table_index + + for name, ad in arm_data.items(): + if name.endswith("__disp"): + continue + scores = np.dot(q_emb[i], ad["emb"].T) + order = top_k_with_maxpool(scores, ad["chunk_ids"], TOP_K_MAX, ad["kind"] == "eu") + page_at = page_hit_at_k(order, ad["pages"], gold_page) + table_at = table_hit_at_k(order, ad["table_idx"], gold_table_index) if gold_table_index is not None else None + top1_idx = order[0] if order else None + disp = arm_data.get(name + "__disp") + top1_text = (disp[top1_idx] if disp else ad["texts"][top1_idx]) if top1_idx is not None else "" + row["arms"][name] = { + "page_hit_at": page_at, + "table_hit_at": table_at, + "evidence_hit_at_1": evidence_hit(qa.get("answer_spec"), top1_text), + "order": order, "texts_ref": name, + } + + if stats["n_tables"] == 0 and "base" in row["arms"] and "eu" in row["arms"]: + b_arm, e_arm = row["arms"]["base"], row["arms"]["eu"] + if b_arm["page_hit_at"] != e_arm["page_hit_at"] or b_arm["evidence_hit_at_1"] != e_arm["evidence_hit_at_1"]: + print(f" [회귀경고] {doc_id}/{qa.get('qid','?')}: 표 0개인데 base/eu 판정 불일치 " + f"(base.page_hit={b_arm['page_hit_at']} eu.page_hit={e_arm['page_hit_at']} " + f"base.ev={b_arm['evidence_hit_at_1']} eu.ev={e_arm['evidence_hit_at_1']})") + + rows.append(row) + + n = len(rows) + b1 = sum(r["arms"]["base"]["page_hit_at"][1] for r in rows) / n + e1 = sum(r["arms"]["eu"]["page_hit_at"][1] for r in rows) / n + print(f" [{n}문항] PageHit@1 base={b1:.3f} eu={e1:.3f}") + + gc.collect() + return rows, arm_data + + +# =========================================================================== +# 7. 실행 / 집계 +# =========================================================================== + +def run(pdf_dir: Path, qa_dir: Path, out_dir: Path, dev_only: bool, max_pdfs, + gemini_api_key: Optional[str], llm_sample_n: int) -> dict: + """전체 벤치마크를 실행하고 summary dict를 반환(+ 결과 JSON 저장).""" + import evidence_chunker + from sentence_transformers import SentenceTransformer + import torch + + device = "cuda" if torch.cuda.is_available() else "cpu" + tag = "dev20" if dev_only else "full90" + print(f"[setup] device={device} scope={tag} evidence_chunker={evidence_chunker.__version__}") + + model = SentenceTransformer(EMBED_MODEL_NAME, device=device) + + tokenizer = None + try: + import tiktoken + tokenizer = tiktoken.get_encoding("cl100k_base") + except Exception: + print("[warn] tiktoken 없음 — 토큰 예산을 단어수로 근사") + + pairs = pdf_qa_pairs(pdf_dir, qa_dir, dev_only, max_pdfs) + if not pairs: + raise SystemExit("[ERR] PDF-QA 쌍 없음 — --pdf-dir/--qa-dir 확인") + + all_rows, all_arm_data = [], {} + for i, (pdf, qa, doc_id) in enumerate(pairs, 1): + print(f"\n[{i}/{len(pairs)}]", end="") + result = evaluate_document(pdf, qa, doc_id, model, tokenizer) + if not result: + continue + rows, arm_data = result + all_rows += rows + all_arm_data[doc_id] = arm_data + + print(f"\n{'='*100}\n HEADLINE — {len(all_rows)}문항 ({tag})\n{'='*100}") + + ARMS = ["base", "eu", "eu_no_caption", "eu_no_context", "eu_no_both"] + for k in K_LIST: + line = " ".join( + f"{a}={sum(r['arms'][a]['page_hit_at'][k] for r in all_rows if a in r['arms']) / max(1, sum(1 for r in all_rows if a in r['arms'])):.3f}" + for a in ARMS + ) + print(f" PageHit@{k}: {line}") + ev_line = " ".join( + f"{a}={sum(r['arms'][a]['evidence_hit_at_1'] for r in all_rows if a in r['arms']) / max(1, sum(1 for r in all_rows if a in r['arms'])):.3f}" + for a in ARMS + ) + print(f" evidence_hit@1: {ev_line}") + + table_rows = [r for r in all_rows if r.get("gold_table_index") is not None] + if table_rows: + for k in K_LIST: + line = " ".join( + f"{a}={sum(r['arms'][a]['table_hit_at'][k] for r in table_rows if a in r['arms'] and r['arms'][a].get('table_hit_at')) / max(1, sum(1 for r in table_rows if a in r['arms'] and r['arms'][a].get('table_hit_at'))):.3f}" + for a in ARMS + ) + print(f" TableHit@{k}(표 질문 {len(table_rows)}개만): {line}") + else: + print(" TableHit@k: 표 질문(meta.table_index 존재) 없음 — 건너뜀") + + doc_ids = [r["doc_id"] for r in all_rows if "base" in r["arms"] and "eu" in r["arms"]] + b_vec = [int(r["arms"]["base"]["page_hit_at"][1]) for r in all_rows if "base" in r["arms"] and "eu" in r["arms"]] + e_vec = [int(r["arms"]["eu"]["page_hit_at"][1]) for r in all_rows if "base" in r["arms"] and "eu" in r["arms"]] + stats_page1 = paired_statistics(doc_ids, b_vec, e_vec) if doc_ids else None + + ablation_stats = {} + for arm in ("eu_no_caption", "eu_no_context", "eu_no_both"): + dids = [r["doc_id"] for r in all_rows if arm in r["arms"] and "eu" in r["arms"]] + full_vec = [int(r["arms"]["eu"]["evidence_hit_at_1"]) for r in all_rows if arm in r["arms"] and "eu" in r["arms"]] + ablated_vec = [int(r["arms"][arm]["evidence_hit_at_1"]) for r in all_rows if arm in r["arms"] and "eu" in r["arms"]] + if dids: + ablation_stats[arm] = paired_statistics(dids, ablated_vec, full_vec) # ablated -> full (얼마나 잃었나) + + llm_rows = run_llm_eval(all_rows, all_arm_data, gemini_api_key, llm_sample_n, tokenizer) + llm_summary = None + if llm_rows: + dids = [r["doc_id"] for r in llm_rows] + b_llm = [int(r["baseline_llm_correct"]) for r in llm_rows] + e_llm = [int(r["eu_llm_correct"]) for r in llm_rows] + llm_summary = paired_statistics(dids, b_llm, e_llm) + print(f"\n[LLM 실답변 평가] n={len(llm_rows)} " + f"baseline={llm_summary['baseline']:.3f} eu={llm_summary['treatment']:.3f} " + f"diff={llm_summary['difference_pp']:+.1f}pp") + + def _blk(rs: list) -> Optional[dict]: + """부분집합 요약 — base/eu 둘 다 페이지 일치 기준(page_hit_at[1])으로 통일.""" + rs = [r for r in rs if "base" in r["arms"] and "eu" in r["arms"]] + if not rs: + return None + n = len(rs) + return { + "n": n, + "baseline_recall_page": round(sum(r["arms"]["base"]["page_hit_at"][1] for r in rs) / n, 4), + "eu_recall_page": round(sum(r["arms"]["eu"]["page_hit_at"][1] for r in rs) / n, 4), + "baseline_evidence_hit": round(sum(r["arms"]["base"]["evidence_hit_at_1"] for r in rs) / n, 4), + "eu_evidence_hit": round(sum(r["arms"]["eu"]["evidence_hit_at_1"] for r in rs) / n, 4), + } + + by_type = {t: _blk([r for r in all_rows if r["type"] == t]) + for t in ("cell_value", "table_about", "context_dependent")} + + ctx_rows = [r for r in all_rows if r["type"] == "context_dependent"] + context_dependent_slices = { + "in_window": _blk([r for r in ctx_rows if (r["meta"].get("dist_pt") or 0) <= BBOX_THRESHOLD]), + "out_window": _blk([r for r in ctx_rows if (r["meta"].get("dist_pt") or 0) > BBOX_THRESHOLD]), + "explicit_ref": _blk([r for r in ctx_rows if r["meta"].get("ctx_explicit_ref") is True]), + "no_explicit_ref": _blk([r for r in ctx_rows if r["meta"].get("ctx_explicit_ref") is False]), + } if ctx_rows else {} + + meta_slices = { + "cross_table": _blk([r for r in all_rows if (r["meta"].get("n_tables_on_page") or 0) >= 2]), + "single_table": _blk([r for r in all_rows if (r["meta"].get("n_tables_on_page") or 0) == 1]), + "toc_zone": _blk([r for r in all_rows if (r["meta"].get("page_index") or 1) <= 0.1]), + "multi_header": _blk([r for r in all_rows if (r["meta"].get("header_rows") or 1) > 1]), + } + + summary = { + "config": {"scope": tag, "bbox_threshold": BBOX_THRESHOLD, "sim_threshold": SIM_THRESHOLD, + "embed_model": EMBED_MODEL_NAME, "context_token_budget": CONTEXT_TOKEN_BUDGET, + "llm_sample_n": len(llm_rows), "evidence_chunker": evidence_chunker.__version__}, + "n_questions": len(all_rows), + "page_hit_at_k": {k: {a: round(sum(r["arms"][a]["page_hit_at"][k] for r in all_rows if a in r["arms"]) / + max(1, sum(1 for r in all_rows if a in r["arms"])), 4) + for a in ARMS} for k in K_LIST}, + "evidence_hit_at_1": {a: round(sum(r["arms"][a]["evidence_hit_at_1"] for r in all_rows if a in r["arms"]) / + max(1, sum(1 for r in all_rows if a in r["arms"])), 4) for a in ARMS}, + "table_hit_at_k": {k: {a: round(sum(r["arms"][a]["table_hit_at"][k] for r in table_rows if a in r["arms"] and r["arms"][a].get("table_hit_at")) / + max(1, sum(1 for r in table_rows if a in r["arms"] and r["arms"][a].get("table_hit_at"))), 4) + for a in ARMS} for k in K_LIST} if table_rows else None, + "n_table_questions": len(table_rows), + "by_type": by_type, + "context_dependent_slices": context_dependent_slices, + "meta_slices": meta_slices, + "paired_statistics_page_hit_1_base_vs_eu": stats_page1, + "ablation_paired_statistics": ablation_stats, + "llm_answer_eval": {"summary": llm_summary, "rows": llm_rows}, + } + out_dir.mkdir(parents=True, exist_ok=True) + (out_dir / f"bench_v6_{tag}.json").write_text(json.dumps(summary, indent=2, ensure_ascii=False), encoding="utf-8") + (out_dir / f"bench_v6_{tag}_rows.json").write_text( + json.dumps([{**r, "arms": {a: {k: v for k, v in d.items() if k != "order"} for a, d in r["arms"].items()}} + for r in all_rows], indent=2, ensure_ascii=False), encoding="utf-8") + print(f"\n저장: {out_dir}/bench_v6_{tag}.json") + return summary + + +# =========================================================================== +# 8. CLI +# =========================================================================== + +def parse_args() -> argparse.Namespace: + p = argparse.ArgumentParser(description=__doc__.strip().splitlines()[0]) + p.add_argument("--pdf-dir", type=Path, required=True, help="벤치마크 PDF 디렉토리") + p.add_argument("--qa-dir", type=Path, required=True, + help="generate_qa_docling.py가 만든 {문서}_qa.json이 있는 디렉토리") + p.add_argument("--out-dir", type=Path, required=True, help="결과 JSON을 저장할 디렉토리") + p.add_argument("--dev-only", action="store_true", help="dev 서브셋(20개 문서)만 실행") + p.add_argument("--max-pdfs", type=int, default=None, help="추가 상한 (디버깅용)") + p.add_argument("--gemini-api-key", type=str, default=None, + help="LLM 실답변 평가용. 미지정 시 GEMINI_API_KEY 환경변수 사용, " + "둘 다 없으면 그 단계만 건너뜀") + p.add_argument("--llm-sample-n", type=int, default=LLM_SAMPLE_N, + help="LLM 실답변 평가 샘플 문항 수(유형별 층화)") + p.add_argument("--bbox-threshold", type=float, default=None, help="BBOX_THRESHOLD") + p.add_argument("--sim-threshold", type=float, default=None, help="SIM_THRESHOLD") + p.add_argument("--embed-model", type=str, default=None, help="EMBED_MODEL_NAME") + p.add_argument("--encode-batch", type=int, default=None, help="ENCODE_BATCH") + p.add_argument("--context-token-budget", type=int, default=None, help="CONTEXT_TOKEN_BUDGET") + return p.parse_args() + + +def main() -> None: + global BBOX_THRESHOLD, SIM_THRESHOLD, EMBED_MODEL_NAME, ENCODE_BATCH, CONTEXT_TOKEN_BUDGET + + args = parse_args() + + if args.bbox_threshold is not None: + BBOX_THRESHOLD = args.bbox_threshold + if args.sim_threshold is not None: + SIM_THRESHOLD = args.sim_threshold + if args.embed_model is not None: + EMBED_MODEL_NAME = args.embed_model + if args.encode_batch is not None: + ENCODE_BATCH = args.encode_batch + if args.context_token_budget is not None: + CONTEXT_TOKEN_BUDGET = args.context_token_budget + + gemini_api_key = args.gemini_api_key or os.environ.get("GEMINI_API_KEY") + + run(args.pdf_dir, args.qa_dir, args.out_dir, args.dev_only, args.max_pdfs, + gemini_api_key, args.llm_sample_n) + + +if __name__ == "__main__": + main() diff --git a/benchmarks/measure_recall.py b/benchmarks/measure_recall.py index 1f7df9d..6d89045 100644 --- a/benchmarks/measure_recall.py +++ b/benchmarks/measure_recall.py @@ -1,27 +1,13 @@ """ -measure_recall.py — baseline(Docling HybridChunker) vs EvidenceChunker 벤치마크 +measure_recall.py — CI 벤치마크 러너. 평가 로직의 단일 소스는 benchmark.py이며, +이 파일은 benchmark.run()을 호출하고 DagsHub/MLflow 로깅만 담당하는 얇은 +래퍼다. 채점 로직을 고치려면 benchmark.py만 고치면 된다. -같은 PDF+질문셋에서 두 코퍼스를 만들어 Recall@1 / EM을 비교한다. +사용법(GitHub Actions에서 실제로 쓰는 형태): + python benchmarks/measure_recall.py \\ + --pdf-dir ./data/pdfs --qa-dir ./data/auto_qa --out-dir ./results \\ + --mlflow --dev-only [--max-pdfs 5] - baseline : HybridChunker().chunk(doc) 전체 청크 (표 청크 포함, 원래 방식) - EU : EvidenceChunker.build_corpus() — EvidenceUnit + 비표 본문(카니발라이제이션 제거 후) -' -지표: - Recall@1 (strict) : top-1이 EU 유닛이고 페이지가 정답과 일치 - Recall@1 (page) : top-1 페이지만 일치 (hybrid 청크여도 인정) - EM : top-1 청크 텍스트에 answer_spec이 만족되는가 - -EM이 핵심 지표다. context_dependent는 표와 문단이 같은 페이지라 페이지 기준 -Recall로는 변별이 안 되고, answer_spec의 비대칭 설계(value=표에만, -context_keys=문단에만) 때문에 EU만 구조적으로 통과 가능하다. - -주의: normalize_for_em/has_token/em_hit은 benchmarks/generate_qa_docling.py의 -동일 함수와 반드시 같은 규칙을 유지해야 한다 — 어긋나면 answer_spec이 의미를 잃는다. - -통계: _ci()는 참고용 단순 근사치. 유의성 판단은 paired_statistics()의 문서 단위 -클러스터 부트스트랩 CI(주 지표) + McNemar 검정(보조 지표, 질문 간 독립 가정)을 쓴다. - -사용법: python measure_recall.py --pdf-dir ./data/pdfs --qa-dir ./auto_qa --out-dir ./results python measure_recall.py --pdf-dir ./data/pdfs --qa-dir ./auto_qa --out-dir ./results --dev-only """ @@ -29,704 +15,108 @@ from __future__ import annotations import argparse -import gc -import json -import math import os -import re -from collections import Counter, defaultdict +import sys from pathlib import Path -import numpy as np - -# --------------------------------------------------------------------------- -# 하이퍼파라미터 — 라이브러리 기본값과 동일하게 유지 (sweep은 별도 실험) -# --------------------------------------------------------------------------- - -EMBED_MODEL_NAME = "sentence-transformers/all-MiniLM-L6-v2" -ENCODE_BATCH = 256 -BBOX_THRESHOLD = 300.0 -SIM_THRESHOLD = 0.00 -CTX_WINDOW_PT = 300.0 # dist_pt 슬라이스 기준 (게이트 아님, BBOX_THRESHOLD와 같은 값) - -HEADLINE_TYPES = {"cell_value", "table_about", "context_dependent"} -VALID_SUBSETS = {None, "main"} - -# dev 서브셋 — 수동 QA셋 20개와 동일 문서 구성 -DEV_DOCS = [ - "1. Attention is all you need", "6. DPO", "14. CLIP", "21. T5", "26. MMLU", - "35. APB", "39. risk sharing", "42. rural housing", "45. gao-25-107649", - "47. gao-26-107681", "50. gao-26-107884", "52. gao-26-108011", - "55. gao-26-108116", "60. ieee1", "64. ieee5", "66. ieee7", "70. ieee11", - "72. ieee13", "80. ssrn-1331573", "85. ssrn-2760631", -] - -# 문서 매크로 평균에서 이 문항 수 미만은 제외한다. QA셋은 문서당 1~122문항으로 -# 편차가 커서(median 20), 문항 적은 문서가 매크로 평균에서 과대 대표되는 걸 막는다. -MIN_DOC_N = 10 - - -# =========================================================================== -# 1. 채점 — generate_qa_docling.py §1과 동일 규칙 -# =========================================================================== - -def normalize_for_em(s: str) -> str: - if not s: - return "" - s = s.lower().replace("\xa0", " ") - s = re.sub(r"[^\w.\-±× ]+", " ", s) - s = re.sub(r"\s+", " ", s).strip() - s = re.sub(r"\.$", "", s) - return s - - -def has_token(needle: str, haystack_norm: str) -> bool: - n = normalize_for_em(needle) - if not n: - return False - return re.search(r"(? bool: - """값 + 모든 context_keys가 같은 청크 안에 있으면 정답. - - EU의 'A | B: v'도 baseline의 '| A | v |'도 동일하게 판정된다 — EU 전용 - 통과 경로 없음(생성 단계 CI 불변식으로 전수 검증됨). - """ - if not spec or not chunk_text: - return False - c = normalize_for_em(chunk_text) - if not has_token(spec.get("value", ""), c): - return False - return all(has_token(k, c) for k in spec.get("context_keys", [])) - - -def _legacy_answer_in_chunk(answer: str, chunk_text: str) -> bool: - """answer_spec 없는 구버전 수동 QA셋 호환용 폴백.""" - a, c = normalize_for_em(answer), normalize_for_em(chunk_text) - if not a or not c: - return False - if re.search(r"(?= 4] - if not tokens: - return False - return sum(1 for t in tokens if has_token(t, c)) / len(tokens) >= 0.8 - - -def score_em(qa: dict, chunk_text: str) -> bool: - spec = qa.get("answer_spec") - return em_hit(spec, chunk_text) if spec else _legacy_answer_in_chunk(qa.get("answer", ""), chunk_text) - - -def classify_real_driver(b_ok: bool, e_ok: bool) -> str: - if b_ok and e_ok: - return "both_right" - if b_ok and not e_ok: - return "baseline_win_eu_lose" - if not b_ok and e_ok: - return "eu_win_baseline_lose" - return "both_wrong" - - -# =========================================================================== -# 2. PDF ↔ QA 매칭 -# =========================================================================== - -_LEADING_NUM = re.compile(r"^(\d+)\.") - - -def _sort_key(p: Path): - m = _LEADING_NUM.match(p.name) - return (int(m.group(1)) if m else 10 ** 6, p.name) - - -def pdf_qa_pairs(pdf_dir: Path, qa_dir: Path, dev_only: bool, max_pdfs: int | None) -> list[tuple[Path, Path, str]]: - """(pdf_path, qa_path, doc_id) 목록. auto_qa는 파일명이 PDF와 1:1 대응한다.""" - pairs, unmatched, empty = [], [], [] - for qa_path in sorted(qa_dir.glob("*_qa.json"), key=_sort_key): - if qa_path.name.startswith("_"): - continue - doc_id = qa_path.stem[:-3] if qa_path.stem.endswith("_qa") else qa_path.stem - if dev_only and doc_id not in DEV_DOCS: - continue - pdf_path = pdf_dir / f"{doc_id}.pdf" - if not pdf_path.exists(): - unmatched.append(qa_path.name) - continue - try: - n = len(json.load(open(qa_path, encoding="utf-8"))) - except Exception: - n = 0 - if n == 0: - empty.append(qa_path.name) - continue - pairs.append((pdf_path, qa_path, doc_id)) - - if unmatched: - print(f" [warn] PDF 매칭 실패 {len(unmatched)}건: {unmatched[:5]}") - if empty: - print(f" [skip] 0문항 파일 {len(empty)}건: {empty}") - if dev_only: - missing = [d for d in DEV_DOCS if d not in {p[2] for p in pairs}] - if missing: - print(f" [warn] dev 목록에 있으나 못 찾음: {missing}") - print(f" [pairs] {len(pairs)} PDF+QA") - return pairs[:max_pdfs] if max_pdfs else pairs - - -def chunk_page(c) -> "int | None": - """HybridChunker 청크의 페이지 번호.""" - try: - for di in c.meta.doc_items: - if di.prov: - return di.prov[0].page_no - except Exception: - pass - return None - - -# =========================================================================== -# 3. 코퍼스 구성 — PDF 1회 파싱으로 baseline·EU 동시 생성 -# =========================================================================== -# EvidenceChunker.build_corpus()를 그대로 쓰면 Docling 변환이 2회(baseline용 + -# EU용) 일어난다. 아래는 build_corpus()와 동일한 로직을 doc 재사용 형태로 -# 편 것이며 라이브러리 함수만 호출한다(재구현 아님). PARITY_CHECK=True로 -# 첫 문서에서 build_corpus() 결과와 대조해 동등성을 확인할 수 있다. - -_converter = None - - -def get_converter(): - global _converter - if _converter is None: - from docling.document_converter import DocumentConverter - _converter = DocumentConverter() - return _converter - - -def build_both_corpora(pdf_path: Path, doc_id: str, parity_check: bool = False): - """Returns (baseline: list[HybridChunker chunk], eu_corpus: list[RetrievalChunk], stats: dict).""" - from docling.chunking import HybridChunker - from docling_core.types.doc import DocItemLabel - from evidence_chunker.chunker import build_evidence_units - from evidence_chunker.split import split_oversized_units - from evidence_chunker.parser.docling import DoclingParser - from evidence_chunker.export import TextChunk, filter_consumed_paragraphs - - doc = get_converter().convert(str(pdf_path)).document - - all_chunks = list(HybridChunker().chunk(doc)) - is_table_chunk = lambda c: any(di.label == DocItemLabel.TABLE for di in c.meta.doc_items) - - parsed = DoclingParser().from_doc(doc) - eu_list = build_evidence_units(parsed, BBOX_THRESHOLD, SIM_THRESHOLD, doc_id) - eu_list = split_oversized_units(eu_list) - - non_table = [c for c in all_chunks if not is_table_chunk(c)] - before_dedup = len(non_table) - non_table = filter_consumed_paragraphs(non_table, eu_list) - eu_corpus = eu_list + [TextChunk(c, doc_id, i) for i, c in enumerate(non_table)] - - conf = Counter(eu.caption_confidence for eu in eu_list) - stats = { - "n_tables": len(parsed.tables), - "n_eu": len(eu_list), - "n_split": sum(1 for eu in eu_list if eu.is_split), - "n_baseline_chunks": len(all_chunks), - "n_hybrid_before_dedup": before_dedup, - "n_consumed_removed": before_dedup - len(non_table), - "n_poisoned_captions": sum(1 for eu in eu_list if eu.caption_text and not eu.safe_caption), - "caption_confidence": dict(conf), - } - - if parity_check: - from evidence_chunker import EvidenceChunker - ref = EvidenceChunker(bbox_threshold=BBOX_THRESHOLD, - sim_threshold=SIM_THRESHOLD).build_corpus(pdf_path, doc_id=doc_id) - a = [c.chunk_id for c in eu_corpus] - b = [c.chunk_id for c in ref] - print(f" [parity] build_corpus()와 동일: {a == b} ({len(a)} vs {len(b)})") - - del doc - gc.collect() - return all_chunks, eu_corpus, stats - - -# =========================================================================== -# 4. 문서 1개 평가 -# =========================================================================== - -def evaluate_one(pdf_path: Path, qa_path: Path, doc_id: str, model, device: str, parity_check: bool = False) -> dict: - import numpy as np - - print(f"\n{'-'*62}") - print(f" {doc_id}") - - try: - b_chunks, eu_corpus, stats = build_both_corpora(pdf_path, doc_id, parity_check) - except Exception as e: - print(f" [ERR] {type(e).__name__}: {e}") - return {} - - print(f" 표 {stats['n_tables']} → EU {stats['n_eu']} (분할 {stats['n_split']}) | " - f"baseline 청크 {stats['n_baseline_chunks']}") - print(f" 카니발라이제이션 제거 {stats['n_consumed_removed']}/{stats['n_hybrid_before_dedup']} | " - f"caption {stats['caption_confidence']}") - - b_texts = [c.text for c in b_chunks] - b_pages = [chunk_page(c) for c in b_chunks] - - e_units, e_src, e_pages, e_disp = [], [], [], [] - for c in eu_corpus: - page = c.metadata.get("page_no") - for u in c.retrieval_units: - e_units.append(u) - e_src.append(c.chunk_id) - e_pages.append(page) - e_disp.append(c.text) - - if not b_texts or not e_units: - print(" [ERR] 빈 코퍼스") - return {} - print(f" EU 코퍼스 {len(e_units)} 유닛 (EU {stats['n_eu']} + hybrid " - f"{len(eu_corpus) - stats['n_eu']})") - - enc = lambda xs: model.encode(xs, normalize_embeddings=True, - show_progress_bar=False, batch_size=ENCODE_BATCH) - b_emb, e_emb = enc(b_texts), enc(e_units) - - qa_list = json.load(open(qa_path, encoding="utf-8")) - usable, skipped = [], 0 - for qa in qa_list: - if qa.get("question_delabeled") is None or qa.get("subset", "main") not in VALID_SUBSETS: - skipped += 1 - continue - usable.append(qa) - if not usable: - print(" [warn] 사용 가능 문항 0") - return {} - - q_emb = enc([q["question_delabeled"] for q in usable]) - b_top = np.dot(q_emb, b_emb.T).argmax(axis=1) - e_top = np.dot(q_emb, e_emb.T).argmax(axis=1) - - rows = [] - for i, qa in enumerate(usable): - exp = qa.get("page") - meta = qa.get("meta") or {} - - bi = int(b_top[i]) - b_ok = b_pages[bi] is not None and b_pages[bi] == exp - b_em = score_em(qa, b_texts[bi]) - - ei = int(e_top[i]) - src, epage = e_src[ei], e_pages[ei] - is_eu_unit = "-hybrid-" not in src # TextChunk는 "{doc_id}-hybrid-{i}" - e_ok_strict = is_eu_unit and epage == exp - e_ok_page = epage == exp - e_em = score_em(qa, e_disp[ei]) - - if e_ok_strict: - fail = None - elif not is_eu_unit: - fail = "hybrid won" + (" (page O)" if e_ok_page else " (page X)") - else: - fail = f"EU wrong page (p{epage}, exp p{exp})" - - rows.append({ - "doc_id": doc_id, "qid": qa.get("qid", ""), "type": qa.get("type", "unknown"), - "question": qa["question_delabeled"], "expected_page": exp, - "b_correct": b_ok, "b_em": b_em, - "e_source": src, "e_page": epage, - "e_correct": e_ok_strict, "e_page_correct": e_ok_page, "e_em": e_em, - "fail_reason": fail, - "real_driver": classify_real_driver(b_ok, e_ok_strict), - "e_top1_text": e_units[ei][:120], - "dist_pt": meta.get("dist_pt"), - "ctx_explicit_ref": meta.get("ctx_explicit_ref"), - "n_tables_on_page": meta.get("n_tables_on_page"), - "page_index": meta.get("page_index"), - "header_rows": meta.get("header_rows"), - }) - - n = len(rows) - print(f"\n [{n}문항, 제외 {skipped}]") - print(f" baseline R {sum(r['b_correct'] for r in rows)/n:.3f} EM {sum(r['b_em'] for r in rows)/n:.3f}") - print(f" EU R {sum(r['e_correct'] for r in rows)/n:.3f} EM {sum(r['e_em'] for r in rows)/n:.3f}") - for t in sorted({r["type"] for r in rows}): - s = [r for r in rows if r["type"] == t] - print(f" {t:20s} n={len(s):4d} " - f"R {sum(r['b_correct'] for r in s)/len(s):.3f}->{sum(r['e_correct'] for r in s)/len(s):.3f} " - f"EM {sum(r['b_em'] for r in s)/len(s):.3f}->{sum(r['e_em'] for r in s)/len(s):.3f}") - - gc.collect() - if device == "cuda": - import torch - torch.cuda.empty_cache() - return {"doc_id": doc_id, "rows": rows, "skipped": skipped, **stats} - +# benchmark.py는 같은 디렉토리에 있다. 다른 위치에서 -m 등으로 실행될 경우까지 대비. +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import benchmark as bv6 # noqa: E402 -# =========================================================================== -# 5. 집계 / 리포트 -# =========================================================================== - -def _rate(rows, key): - return (sum(r[key] for r in rows) / len(rows)) if rows else None - - -def _line(label, rows, width=28): - if not rows: - return f" {label:<{width}} n= 0 - -" - n = len(rows) - bR, eR = _rate(rows, "b_correct"), _rate(rows, "e_correct") - bE, eE = _rate(rows, "b_em"), _rate(rows, "e_em") - return (f" {label:<{width}} n={n:5d} " - f"R {bR:.3f}->{eR:.3f} ({(eR-bR)*100:+5.1f}pp) " - f"EM {bE:.3f}->{eE:.3f} ({(eE-bE)*100:+5.1f}pp)") - - -def _ci(n): - """95% CI 반폭(pp) — 표본이 작을 때 관측된 갭이 실제로 유의한지 판단하는 기준. - - 단순 이항분포 최대분산(p=0.5) 근사(Wald)라, 같은 질문에 대한 - baseline/EU 쌍대비교 구조나 문서 내 상관은 반영하지 않는다. 엄밀한 - 유의성 판단에는 아래 paired_statistics()를 쓸 것 — 이 함수는 빠른 - 참고용 오차범위로만 남겨둔다. - """ - return 1.96 * 0.5 / (n ** 0.5) * 100 if n else float("inf") - - -def paired_statistics(doc_ids, baseline, treatment, repeats: int = 50000, seed: int = 20260915) -> dict: - """문서 단위 클러스터 부트스트랩 CI + 보조 McNemar 검정. - - _ci()와 달리 같은 문서에서 나온 질문들을 묶어서(문서를 리샘플링 단위로 - 삼아) 차이의 95% percentile CI를 구한다 — 문서 내 질문 간 상관을 - 반영하는 방식. McNemar는 "baseline만 맞음 vs EU만 맞음"의 비대칭을 - 검정하는 보조 지표로 덧붙이되, 문서 간 의존성은 보정하지 않는다는 - 가정을 명시한다. - - baseline/treatment: 질문별 0/1(또는 bool) 정오답 배열. doc_ids와 길이가 - 같아야 하며, 같은 인덱스가 같은 질문을 가리켜야 한다(쌍대비교 전제). - """ - b = np.asarray(list(baseline), dtype=np.int64) - e = np.asarray(list(treatment), dtype=np.int64) - doc_ids = list(doc_ids) - if len(doc_ids) != len(b) or len(b) != len(e) or not len(b): - raise ValueError("Paired vectors must be nonempty and equal length.") - - grouped: dict = defaultdict(lambda: [0, 0]) # doc_id -> [n_questions, sum(e-b)] - for d, delta in zip(doc_ids, e - b): - grouped[d][0] += 1 - grouped[d][1] += int(delta) - a = np.asarray([grouped[d] for d in sorted(grouped)], dtype=np.int64) - - ci = None - if len(a) >= 2: - rng, values = np.random.default_rng(seed), [] - for start in range(0, repeats, 2048): - idx = rng.integers(0, len(a), size=(min(2048, repeats - start), len(a))) - total = a[idx].sum(axis=1) - values.append(100 * total[:, 1] / total[:, 0]) - ci = np.quantile(np.concatenate(values), [.025, .975]).tolist() - - b_only = int(((b == 1) & (e == 0)).sum()) - e_only = int(((b == 0) & (e == 1)).sum()) - discordant = b_only + e_only - chi2 = max(abs(b_only - e_only) - 1, 0) ** 2 / discordant if discordant else 0.0 - p = math.erfc(math.sqrt(chi2 / 2)) if discordant else 1.0 - - return { - "n_questions": len(b), "n_documents": len(a), - "baseline": float(b.mean()), "treatment": float(e.mean()), - "difference_pp": float(100 * (e - b).mean()), - "cluster_bootstrap_ci95_pp": ci, - "bootstrap": {"unit": "document", "statistic": "micro_rate_difference", - "method": "percentile", "repeats": repeats, "seed": seed}, - "mcnemar_supplementary": { - "method": "chi_square_continuity_corrected", "chi2": chi2, "p_value": p, - "baseline_only": b_only, "treatment_only": e_only, - "assumption": "질문 간 독립 가정 — 문서 내 의존성은 보정하지 않음(보조 지표)", - }, - "ci_note": ("문서를 독립 표집 단위로 취급한 근사치이며, LLM 생성 답변 정확도의 CI가 아님" - if ci is not None else "문서 2개 미만이라 CI 계산 불가"), - } - - -def run(pdf_dir: Path, qa_dir: Path, out_dir: Path, dev_only: bool, max_pdfs: int | None, - parity_check: bool = False, use_mlflow: bool = False) -> None: - from sentence_transformers import SentenceTransformer - import torch - - device = "cuda" if torch.cuda.is_available() else "cpu" - tag = "dev20" if dev_only else "full90" - - if use_mlflow: - import mlflow - if not os.environ.get("MLFLOW_TRACKING_URI"): - mlflow_db = (Path.cwd() / "mlflow.db").resolve() - mlflow.set_tracking_uri(f"sqlite:///{mlflow_db.as_posix()}") - mlflow.set_experiment("evidence-chunker-benchmark") - mlflow.start_run(run_name=tag) - mlflow.log_params({ - "tag": tag, - "bbox_threshold": BBOX_THRESHOLD, - "sim_threshold": SIM_THRESHOLD, - "embed_model": EMBED_MODEL_NAME, - "encode_batch": ENCODE_BATCH, - "min_doc_n": MIN_DOC_N, - "max_pdfs": max_pdfs, - }) - - pairs = pdf_qa_pairs(pdf_dir, qa_dir, dev_only, max_pdfs) - if not pairs: - print("[ERR] PDF-QA 쌍 없음") - if use_mlflow: - mlflow.end_run(status="FAILED") - return - - import evidence_chunker - print(f"[setup] evidence_chunker {evidence_chunker.__version__} device={device}") - print(f"[setup] scope={'dev 20' if dev_only else 'full 90'}") - print(f"[setup] bbox={BBOX_THRESHOLD}pt sim={SIM_THRESHOLD} batch={ENCODE_BATCH}") - print(f"\n[model] {EMBED_MODEL_NAME} on {device}") - model = SentenceTransformer(EMBED_MODEL_NAME, device=device) - - results, rows = [], [] - for i, (pdf, qa, doc_id) in enumerate(pairs, 1): - print(f"\n[{i}/{len(pairs)}]", end="") - r = evaluate_one(pdf, qa, doc_id, model, device, parity_check) - if r: - results.append(r) - rows += r["rows"] - - if not rows: - print("[ERR] 결과 없음") - if use_mlflow: - mlflow.end_run(status="FAILED") - return - - N = len(rows) - print(f"\n{'='*100}") - print(f" HEADLINE — {len(results)} PDF, {N} 문항 ({'dev 20' if dev_only else 'full 90'})") - print(f"{'='*100}") - print(_line("전체", rows)) - - print(f"\n ── 유형별 ──") - for t in ("cell_value", "table_about", "context_dependent"): - print(_line(t, [r for r in rows if r["type"] == t])) - - ctx = [r for r in rows if r["type"] == "context_dependent"] - if ctx: - print(f"\n ── context_dependent 상세 ──") - inw = [r for r in ctx if r["dist_pt"] is not None and r["dist_pt"] <= CTX_WINDOW_PT] - outw = [r for r in ctx if r["dist_pt"] is not None and r["dist_pt"] > CTX_WINDOW_PT] - print(_line(f"dist <= {CTX_WINDOW_PT:.0f}pt (창 안)", inw)) - print(_line(f"dist > {CTX_WINDOW_PT:.0f}pt (창 밖)", outw)) - if outw: - gap = (_rate(outw, "e_em") - _rate(outw, "b_em")) * 100 - half = _ci(len(outw)) - print(f" 창 밖 n={len(outw)} EM 갭 {gap:+.1f}pp CI ±{half:.1f}pp") - print(_line("explicit_ref=True", [r for r in ctx if r["ctx_explicit_ref"] is True])) - print(_line("explicit_ref=False", [r for r in ctx if r["ctx_explicit_ref"] is False])) - - print(f"\n ── meta 슬라이스 ──") - print(_line("cross-table (표>=2/page)", [r for r in rows if (r["n_tables_on_page"] or 0) >= 2])) - print(_line("single-table (표=1/page)", [r for r in rows if (r["n_tables_on_page"] or 0) == 1])) - print(_line("ToC 구간 (page_index<=.1)", - [r for r in rows if r["page_index"] is not None and r["page_index"] <= 0.1])) - print(_line("다단 헤더 (header_rows>1)", [r for r in rows if (r["header_rows"] or 1) > 1])) - - # 문서 매크로 평균 — 문항 수 편차(1~122문항)로 인한 왜곡을 micro 지표와 함께 교차 확인 - by_doc = defaultdict(list) - for r in rows: - by_doc[r["doc_id"]].append(r) - - def _macro(docs): - if not docs: - return None - n = len(docs) - return {k: sum(_rate(v, k) for v in docs) / n - for k in ("b_correct", "e_correct", "b_em", "e_em")} - - all_docs = list(by_doc.values()) - big_docs = [v for v in all_docs if len(v) >= MIN_DOC_N] - ma, mbig = _macro(all_docs), _macro(big_docs) - - def _mline(label, m, n): - if not m: - return f" {label:<28} n={n:5d} -" - return (f" {label:<28} n={n:5d} " - f"R {m['b_correct']:.3f}->{m['e_correct']:.3f} " - f"({(m['e_correct']-m['b_correct'])*100:+5.1f}pp) " - f"EM {m['b_em']:.3f}->{m['e_em']:.3f} " - f"({(m['e_em']-m['b_em'])*100:+5.1f}pp)") - - print(f"\n ── 문서 매크로 평균 ──") - print(_mline("전체 문서", ma, len(all_docs))) - print(_mline(f"{MIN_DOC_N}문항 이상만", mbig, len(big_docs))) - mb, me = ma["b_correct"], ma["e_correct"] - mbe, mee = ma["b_em"], ma["e_em"] - - rd = Counter(r["real_driver"] for r in rows) - print(f"\n [real_driver] both_right={rd['both_right']} " - f"baseline_win_eu_lose={rd['baseline_win_eu_lose']} " - f"eu_win_baseline_lose={rd['eu_win_baseline_lose']} both_wrong={rd['both_wrong']}") - - paired_em = paired_statistics([r["doc_id"] for r in rows], - [r["b_em"] for r in rows], [r["e_em"] for r in rows]) - ci = paired_em["cluster_bootstrap_ci95_pp"] - ci_str = f"[{ci[0]:+.1f}, {ci[1]:+.1f}]pp" if ci else "N/A(문서<2)" - mc = paired_em["mcnemar_supplementary"] - print(f"\n [EM 통계 검정] 문서 단위 부트스트랩 차이 {paired_em['difference_pp']:+.1f}pp " - f"95% CI {ci_str}") - print(f" [McNemar 보조] baseline만 정답 {mc['baseline_only']} EU만 정답 {mc['treatment_only']} " - f"chi2={mc['chi2']:.2f} p={mc['p_value']:.2e}") - - fails = Counter(r["fail_reason"] for r in rows if r["fail_reason"]) - if fails: - print(f"\n EU 실패 사유 상위") - for k, v in fails.most_common(8): - print(f" {v:5d} {k}") - - tc = sum(r["n_consumed_removed"] for r in results) - tb = sum(r["n_hybrid_before_dedup"] for r in results) - tp = sum(r["n_poisoned_captions"] for r in results) - ts = sum(r["n_split"] for r in results) - conf = Counter() - for r in results: - conf.update(r["caption_confidence"]) - ca = sum(conf.values()) or 1 - print(f"\n [파이프라인] EU {sum(r['n_eu'] for r in results)} (분할 {ts}) " - f"dedup {tc}/{tb} figure-caption 필터 {tp}") - print(f" [caption_confidence] direct {conf['direct']} ({conf['direct']/ca:.0%}) " - f"inferred {conf['inferred']} ({conf['inferred']/ca:.0%}) " - f"none {conf['none']} ({conf['none']/ca:.0%})") - - print(f"\n{'-'*100}") - print(f" {'문서':40s} {'N':>5} {'base_R':>7} {'EU_R':>7} {'base_EM':>8} {'EU_EM':>7}") - print(f"{'-'*100}") - for d, rs in sorted(by_doc.items(), - key=lambda kv: _rate(kv[1], "e_correct") - _rate(kv[1], "b_correct")): - bR, eR = _rate(rs, "b_correct"), _rate(rs, "e_correct") - flag = " *" if eR > bR else (" " if abs(eR - bR) < 1e-9 else " v") - print(f" {d[:40]:40s} {len(rs):5d} {bR:7.3f} {eR:7.3f} " - f"{_rate(rs,'b_em'):8.3f} {_rate(rs,'e_em'):7.3f}{flag}") - - def blk(rs): - if not rs: - return None - return {"n": len(rs), - "baseline_recall": round(_rate(rs, "b_correct"), 4), - "eu_recall": round(_rate(rs, "e_correct"), 4), - "baseline_em": round(_rate(rs, "b_em"), 4), - "eu_em": round(_rate(rs, "e_em"), 4), - "ci_halfwidth_pp": round(_ci(len(rs)), 2)} - - summary = { - "config": {"scope": tag, "qa_dir": str(qa_dir), "embed_model": EMBED_MODEL_NAME, - "bbox_threshold": BBOX_THRESHOLD, "sim_threshold": SIM_THRESHOLD, - "scoring": "answer_spec / em_hit", "evidence_chunker": evidence_chunker.__version__}, - "overall": blk(rows), - "by_type": {t: blk([r for r in rows if r["type"] == t]) - for t in ("cell_value", "table_about", "context_dependent")}, - "context_dependent_slices": { - "in_window": blk([r for r in ctx if r["dist_pt"] is not None and r["dist_pt"] <= CTX_WINDOW_PT]), - "out_window": blk([r for r in ctx if r["dist_pt"] is not None and r["dist_pt"] > CTX_WINDOW_PT]), - "explicit_ref": blk([r for r in ctx if r["ctx_explicit_ref"] is True]), - "no_explicit_ref": blk([r for r in ctx if r["ctx_explicit_ref"] is False]), - } if ctx else {}, - "meta_slices": { - "cross_table": blk([r for r in rows if (r["n_tables_on_page"] or 0) >= 2]), - "single_table": blk([r for r in rows if (r["n_tables_on_page"] or 0) == 1]), - "toc_zone": blk([r for r in rows if r["page_index"] is not None and r["page_index"] <= 0.1]), - "multi_header": blk([r for r in rows if (r["header_rows"] or 1) > 1]), - }, - "macro_average": {"documents": len(all_docs), - "baseline_recall": round(mb, 4), "eu_recall": round(me, 4), - "baseline_em": round(mbe, 4), "eu_em": round(mee, 4)}, - "macro_average_min_n": ({"documents": len(big_docs), "min_questions": MIN_DOC_N, - "baseline_recall": round(mbig["b_correct"], 4), - "eu_recall": round(mbig["e_correct"], 4), - "baseline_em": round(mbig["b_em"], 4), - "eu_em": round(mbig["e_em"], 4)} if mbig else None), - "doc_question_counts": {d: len(v) for d, v in by_doc.items()}, - "real_driver": dict(rd), - "paired_statistics_em": paired_em, - "fail_reasons": dict(fails), - "pipeline": {"n_eu": sum(r["n_eu"] for r in results), "n_split": ts, - "dedup_removed": tc, "hybrid_before_dedup": tb, - "poisoned_captions": tp, "caption_confidence": dict(conf)}, - "by_doc": {d: blk(rs) for d, rs in by_doc.items()}, - } - out_dir.mkdir(parents=True, exist_ok=True) - (out_dir / f"bench_{tag}.json").write_text( - json.dumps(summary, indent=2, ensure_ascii=False), encoding="utf-8") - (out_dir / f"bench_{tag}_rows.json").write_text( - json.dumps(rows, indent=2, ensure_ascii=False), encoding="utf-8") - print(f"\n저장: bench_{tag}.json / bench_{tag}_rows.json ({len(rows)} rows)") - - if use_mlflow: - mlflow.log_metrics(summary["macro_average"]) - if summary.get("macro_average_min_n"): - mlflow.log_metrics({ - f"minN_{k}": v for k, v in summary["macro_average_min_n"].items() - if isinstance(v, (int, float)) - }) - # QA 유형별(cell_value / table_about / context_dependent) 지표도 - # DagsHub에서 바로 비교할 수 있게 별도 metric으로 남긴다. - for t, blk_v in summary.get("by_type", {}).items(): - if not blk_v: - continue - mlflow.log_metrics({ - f"type_{t}_{k}": v for k, v in blk_v.items() - if isinstance(v, (int, float)) - }) - mlflow.log_artifact(str(out_dir / f"bench_{tag}.json")) - mlflow.log_artifact(str(out_dir / f"bench_{tag}_rows.json")) - mlflow.end_run() - - -# =========================================================================== -# 6. CLI -# =========================================================================== def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser(description=__doc__.strip().splitlines()[0]) p.add_argument("--pdf-dir", type=Path, required=True, help="벤치마크 PDF 디렉토리") p.add_argument("--qa-dir", type=Path, required=True, - help="generate_qa_docling.py가 만든 {문서}_qa.json이 있는 디렉토리") + help="generate_qa_docling.py가 만든 {문서}_qa.json이 있는 디렉토리") p.add_argument("--out-dir", type=Path, required=True, help="결과 JSON을 저장할 디렉토리") p.add_argument("--dev-only", action="store_true", help="dev 서브셋(20개 문서)만 실행") - p.add_argument("--max-pdfs", type=int, default=None, help="추가 상한 (디버깅용)") - p.add_argument("--parity-check", action="store_true", - help="첫 문서에서 EvidenceChunker.build_corpus() 결과와 대조") - p.add_argument("--mlflow", action="store_true") - # 스윕 자동화 + p.add_argument("--max-pdfs", type=int, default=None, help="추가 상한 (디버깅/smoke용)") + p.add_argument("--mlflow", action="store_true", help="DagsHub/MLflow에 결과 로깅") + p.add_argument("--gemini-api-key", type=str, default=None, + help="LLM 실답변 평가용. 미지정 시 GEMINI_API_KEY 환경변수 사용, " + "둘 다 없으면 그 단계만 건너뜀(CI 기본 상태)") + p.add_argument("--llm-sample-n", type=int, default=bv6.LLM_SAMPLE_N, + help="LLM 실답변 평가 샘플 문항 수(유형별 층화)") p.add_argument("--bbox-threshold", type=float, default=None, help="BBOX_THRESHOLD") p.add_argument("--sim-threshold", type=float, default=None, help="SIM_THRESHOLD") p.add_argument("--embed-model", type=str, default=None, help="EMBED_MODEL_NAME") p.add_argument("--encode-batch", type=int, default=None, help="ENCODE_BATCH") + p.add_argument("--context-token-budget", type=int, default=None, help="CONTEXT_TOKEN_BUDGET") return p.parse_args() def main() -> None: - global BBOX_THRESHOLD, SIM_THRESHOLD, EMBED_MODEL_NAME, ENCODE_BATCH, CTX_WINDOW_PT - args = parse_args() if args.bbox_threshold is not None: - BBOX_THRESHOLD = args.bbox_threshold - CTX_WINDOW_PT = args.bbox_threshold # 두 값은 항상 같이 움직임 + bv6.BBOX_THRESHOLD = args.bbox_threshold if args.sim_threshold is not None: - SIM_THRESHOLD = args.sim_threshold + bv6.SIM_THRESHOLD = args.sim_threshold if args.embed_model is not None: - EMBED_MODEL_NAME = args.embed_model + bv6.EMBED_MODEL_NAME = args.embed_model if args.encode_batch is not None: - ENCODE_BATCH = args.encode_batch + bv6.ENCODE_BATCH = args.encode_batch + if args.context_token_budget is not None: + bv6.CONTEXT_TOKEN_BUDGET = args.context_token_budget + + gemini_api_key = args.gemini_api_key or os.environ.get("GEMINI_API_KEY") + tag = "dev20" if args.dev_only else "full90" - run(args.pdf_dir, args.qa_dir, args.out_dir, args.dev_only, args.max_pdfs, args.parity_check, - args.mlflow) + mlflow = None + if args.mlflow: + import mlflow as _mlflow + mlflow = _mlflow + if not os.environ.get("MLFLOW_TRACKING_URI"): + mlflow_db = (Path.cwd() / "mlflow.db").resolve() + mlflow.set_tracking_uri(f"sqlite:///{mlflow_db.as_posix()}") + mlflow.set_experiment("evidence-chunker-benchmark") + mlflow.start_run(run_name=tag) + mlflow.log_params({ + "tag": tag, + "bbox_threshold": bv6.BBOX_THRESHOLD, + "sim_threshold": bv6.SIM_THRESHOLD, + "embed_model": bv6.EMBED_MODEL_NAME, + "encode_batch": bv6.ENCODE_BATCH, + "max_pdfs": args.max_pdfs, + "engine": "benchmark.py", + }) + + try: + summary = bv6.run(args.pdf_dir, args.qa_dir, args.out_dir, args.dev_only, args.max_pdfs, + gemini_api_key, args.llm_sample_n) + except Exception: + if mlflow is not None: + mlflow.end_run(status="FAILED") + raise + + if mlflow is not None: + # arm별 dict를 "지표_arm" 형태로 평탄화해서 MLflow에 로깅 + flat = {} + for k, arms in summary["page_hit_at_k"].items(): + flat.update({f"page_hit_at_{k}_{a}": v for a, v in arms.items()}) + flat.update({f"evidence_hit_at_1_{a}": v for a, v in summary["evidence_hit_at_1"].items()}) + if summary.get("table_hit_at_k"): + for k, arms in summary["table_hit_at_k"].items(): + flat.update({f"table_hit_at_{k}_{a}": v for a, v in arms.items()}) + flat["n_questions"] = summary["n_questions"] + flat["n_table_questions"] = summary["n_table_questions"] + if summary.get("paired_statistics_page_hit_1_base_vs_eu"): + flat["page_hit_1_diff_pp"] = summary["paired_statistics_page_hit_1_base_vs_eu"]["difference_pp"] + if summary.get("llm_answer_eval", {}).get("summary"): + llm = summary["llm_answer_eval"]["summary"] + flat["llm_baseline"] = llm["baseline"] + flat["llm_eu"] = llm["treatment"] + flat["llm_diff_pp"] = llm["difference_pp"] + for t, blk in summary.get("by_type", {}).items(): + if not blk: + continue + flat.update({f"type_{t}_{k}": v for k, v in blk.items() if isinstance(v, (int, float))}) + + mlflow.log_metrics(flat) + mlflow.log_artifact(str(args.out_dir / f"bench_v6_{tag}.json")) + mlflow.log_artifact(str(args.out_dir / f"bench_v6_{tag}_rows.json")) + mlflow.end_run() if __name__ == "__main__": diff --git a/scripts/summarize_results.py b/scripts/summarize_results.py index a13bdc5..332e1c1 100644 --- a/scripts/summarize_results.py +++ b/scripts/summarize_results.py @@ -1,7 +1,12 @@ """ -summarize_results.py — results/bench_*.json 을 읽어서 GitHub Actions 의 +summarize_results.py — results/bench_v6_*.json 을 읽어서 GitHub Actions 의 Step Summary(Actions 탭에서 바로 보이는 마크다운)로 출력한다. +benchmark.py의 summary 스키마(overall 없음, page_hit_at_k/evidence_hit_at_1/ +table_hit_at_k/by_type 등)를 기준으로 한다 — measure_recall.py가 옛 스키마를 +쓰던 시절 이 스크립트도 옛 필드명(overall/baseline_recall/eu_recall)을 +읽었는데, 엔진이 benchmark.py로 통합되며 같이 옮겨졌다. + 사용법 (CI 안에서): python scripts/summarize_results.py --label "dev-only (20개)" """ @@ -16,7 +21,7 @@ def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser(description=__doc__.strip().splitlines()[0]) - p.add_argument("--results-dir", default="results", help="bench_*.json 이 있는 디렉토리") + p.add_argument("--results-dir", default="results", help="bench_v6_*.json 이 있는 디렉토리") p.add_argument("--label", default="", help="이번 실행 범위 라벨 (smoke/dev/full 등)") return p.parse_args() @@ -37,29 +42,50 @@ def main() -> None: for f in files: data = json.load(open(f, encoding="utf-8")) - overall = data.get("overall") or {} + n = data.get("n_questions", "-") + page_hit_1 = (data.get("page_hit_at_k") or {}).get("1") or (data.get("page_hit_at_k") or {}).get(1) or {} + evidence_hit = data.get("evidence_hit_at_1") or {} + lines.append(f"\n## {os.path.basename(f)}\n") lines.append("| 지표 | baseline | EU | 문항 수 |") lines.append("|---|---|---|---|") lines.append( - f"| Recall | {overall.get('baseline_recall', '-')} | " - f"{overall.get('eu_recall', '-')} | {overall.get('n', '-')} |" + f"| PageHit@1 | {page_hit_1.get('base', '-')} | " + f"{page_hit_1.get('eu', '-')} | {n} |" ) lines.append( - f"| EM | {overall.get('baseline_em', '-')} | " - f"{overall.get('eu_em', '-')} | {overall.get('n', '-')} |" + f"| evidence_hit@1 | {evidence_hit.get('base', '-')} | " + f"{evidence_hit.get('eu', '-')} | {n} |" ) + table_hit_1 = None + if data.get("table_hit_at_k"): + table_hit_1 = (data["table_hit_at_k"].get("1") or data["table_hit_at_k"].get(1)) + if table_hit_1: + lines.append( + f"| TableHit@1 | {table_hit_1.get('base', '-')} | " + f"{table_hit_1.get('eu', '-')} | {data.get('n_table_questions', '-')} |" + ) + + stats = data.get("paired_statistics_page_hit_1_base_vs_eu") + if stats: + ci = stats.get("cluster_bootstrap_ci95_pp") + ci_str = f"[{ci[0]:+.1f}, {ci[1]:+.1f}]pp" if ci else "N/A" + lines.append( + f"\nPageHit@1 문서단위 부트스트랩 차이: **{stats['difference_pp']:+.1f}pp** " + f"(95% CI {ci_str})\n" + ) + by_type = data.get("by_type") or {} type_blocks = {t: b for t, b in by_type.items() if b} if type_blocks: - lines.append("\n**유형별 EM (baseline → EU)**\n") - lines.append("| 유형 | baseline EM | EU EM | n |") + lines.append("\n**유형별 evidence_hit (baseline → EU)**\n") + lines.append("| 유형 | baseline | EU | n |") lines.append("|---|---|---|---|") for t, blk in type_blocks.items(): lines.append( - f"| {t} | {blk.get('baseline_em', '-')} | " - f"{blk.get('eu_em', '-')} | {blk.get('n', '-')} |" + f"| {t} | {blk.get('baseline_evidence_hit', '-')} | " + f"{blk.get('eu_evidence_hit', '-')} | {blk.get('n', '-')} |" ) text = "\n".join(lines) + "\n"