| |
| """Plot signature-to-background accuracy from generated validation JSONL files.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import csv |
| import json |
| import re |
| import textwrap |
| from collections import Counter |
| from dataclasses import dataclass |
| from pathlib import Path |
|
|
| import matplotlib |
|
|
| matplotlib.use("Agg") |
| import matplotlib.pyplot as plt |
|
|
|
|
| SECTIONS = ("dominant", "irreducible", "reducible") |
| STEP_RE = re.compile(r"step(\d+)") |
| HEADER_RE = re.compile(r"^\s*(dominant|irreducible|reducible)\s*:\s*$", re.IGNORECASE) |
| THINK_BLOCK_RE = re.compile(r"(?is)<think>(.*?)</think>") |
| ANSWER_BLOCK_RE = re.compile(r"(?is)<answer>(.*?)</answer>") |
|
|
|
|
| @dataclass |
| class Catalog: |
| id_to_label: dict[str, str] |
| normalized_label_to_id: dict[str, str] |
|
|
|
|
| @dataclass |
| class ParsedResponse: |
| raw_sections: dict[str, list[str]] |
| canonical_sections: dict[str, set[str]] |
| has_think: bool |
| has_answer: bool |
| think: str |
| answer: str |
|
|
|
|
| def normalize_text(value: str) -> str: |
| value = value.strip().lower().replace("->", " to ").replace("→", " to ") |
| value = re.sub(r"^[\-*\d.)\s]+", "", value) |
| value = re.sub(r"[`\"']", "", value) |
| value = re.sub(r"\s+", " ", value) |
| return value.strip(" .;:") |
|
|
|
|
| def load_catalog(path: Path) -> Catalog: |
| payload = json.loads(path.read_text()) |
| id_to_label = {str(item["id"]): str(item["label"]) for item in payload["processes"]} |
| normalized_label_to_id = {normalize_text(label): process_id for process_id, label in id_to_label.items()} |
| return Catalog(id_to_label=id_to_label, normalized_label_to_id=normalized_label_to_id) |
|
|
|
|
| def canonicalize(item: str, catalog: Catalog) -> str: |
| cleaned = item.strip() |
| if cleaned in catalog.id_to_label: |
| return cleaned |
| normalized = normalize_text(cleaned) |
| return catalog.normalized_label_to_id.get(normalized, normalized) |
|
|
|
|
| def display_item(item: str, catalog: Catalog) -> str: |
| return catalog.id_to_label.get(item, item) |
|
|
|
|
| def strip_bullet(line: str) -> str: |
| return re.sub(r"^[\-*\d.)\s]+", "", line.strip()).strip() |
|
|
|
|
| def answer_section_items(answer: str) -> dict[str, list[str]] | None: |
| stripped = answer.strip() |
| if not stripped: |
| return None |
| try: |
| parsed = json.loads(stripped) |
| except json.JSONDecodeError: |
| parsed = None |
| if isinstance(parsed, dict): |
| sections = {section: [] for section in SECTIONS} |
| dominant = parsed.get("dominant") |
| if isinstance(dominant, str) and dominant.strip(): |
| sections["dominant"].append(dominant.strip()) |
| for section in ("irreducible", "reducible"): |
| values = parsed.get(section, []) |
| if isinstance(values, list): |
| sections[section].extend(str(value).strip() for value in values if str(value).strip()) |
| return sections |
|
|
| lines = [strip_bullet(line) for line in stripped.splitlines() if strip_bullet(line)] |
| if len(lines) == 1: |
| return {"dominant": [lines[0]], "irreducible": [], "reducible": []} |
| return None |
|
|
|
|
| def parse_sections(text: str) -> dict[str, list[str]]: |
| sections = {section: [] for section in SECTIONS} |
| current: str | None = None |
| for raw_line in text.splitlines(): |
| line = raw_line.strip() |
| header_match = HEADER_RE.match(line) |
| if header_match: |
| current = header_match.group(1).lower() |
| continue |
| if current and line.startswith("- "): |
| item = strip_bullet(line) |
| if item: |
| sections[current].append(item) |
| return sections |
|
|
|
|
| def parse_response(text: object, catalog: Catalog) -> ParsedResponse: |
| text = str(text or "") |
| think_match = THINK_BLOCK_RE.search(text) |
| answer_match = ANSWER_BLOCK_RE.search(text) |
| sections = None |
| if answer_match: |
| sections = answer_section_items(answer_match.group(1)) |
| if sections is None: |
| sections = parse_sections(text) |
| canonical_sections = { |
| section: {canonicalize(item, catalog) for item in sections[section] if item} |
| for section in SECTIONS |
| } |
| return ParsedResponse( |
| raw_sections=sections, |
| canonical_sections=canonical_sections, |
| has_think=think_match is not None, |
| has_answer=answer_match is not None, |
| think=think_match.group(1).strip() if think_match else "", |
| answer=answer_match.group(1).strip() if answer_match else "", |
| ) |
|
|
|
|
| def is_none_item(item: str, catalog: Catalog) -> bool: |
| return normalize_text(display_item(item, catalog)).startswith("none ") |
|
|
|
|
| def union_backgrounds(parsed: ParsedResponse, catalog: Catalog) -> set[str]: |
| values = set().union(*(parsed.canonical_sections[section] for section in SECTIONS)) |
| return {item for item in values if not is_none_item(item, catalog)} |
|
|
|
|
| def precision_recall_f1(predicted: set[str], expected: set[str]) -> tuple[float, float, float]: |
| if not predicted and not expected: |
| return 1.0, 1.0, 1.0 |
| if not predicted or not expected: |
| return 0.0, 0.0, 0.0 |
| true_positive = len(predicted & expected) |
| precision = true_positive / len(predicted) |
| recall = true_positive / len(expected) |
| f1 = 0.0 if precision + recall == 0 else 2 * precision * recall / (precision + recall) |
| return precision, recall, f1 |
|
|
|
|
| def join_display(items: set[str], catalog: Catalog) -> str: |
| return "; ".join(display_item(item, catalog) for item in sorted(items)) |
|
|
|
|
| def evaluate_row(row: dict, source: Path, catalog: Catalog) -> dict[str, object]: |
| predicted = parse_response(row.get("prediction", ""), catalog) |
| expected = parse_response(row.get("reference", ""), catalog) |
|
|
| pred_dom = predicted.canonical_sections["dominant"] |
| exp_dom = expected.canonical_sections["dominant"] |
| pred_all = union_backgrounds(predicted, catalog) |
| exp_all = union_backgrounds(expected, catalog) |
|
|
| dominant_exact = bool(exp_dom) and pred_dom == exp_dom |
| dominant_hit = bool(exp_dom & pred_dom) |
| dominant_present_anywhere = bool(exp_dom & pred_all) |
| if dominant_exact: |
| category = "dominant exact" |
| elif dominant_hit: |
| category = "dominant plus extra" |
| elif dominant_present_anywhere: |
| category = "right background, wrong section" |
| elif not pred_dom: |
| category = "no dominant parsed" |
| else: |
| category = "dominant missing" |
|
|
| metrics: dict[str, object] = { |
| "source": source.name, |
| "id": row.get("id"), |
| "loss": row.get("loss"), |
| "expected_dominant": join_display(exp_dom, catalog), |
| "predicted_dominant": join_display(pred_dom, catalog), |
| "category": category, |
| "dominant_exact": dominant_exact, |
| "dominant_hit": dominant_hit, |
| "dominant_present_anywhere": dominant_present_anywhere, |
| "all_exact": pred_all == exp_all, |
| "missing_expected_count": len(exp_all - pred_all), |
| "extra_predicted_count": len(pred_all - exp_all), |
| "expected_background_count": len(exp_all), |
| "predicted_background_count": len(pred_all), |
| "prediction_has_think": predicted.has_think, |
| "prediction_has_answer": predicted.has_answer, |
| "reference_has_think": expected.has_think, |
| "reference_has_answer": expected.has_answer, |
| "prediction_think": predicted.think, |
| "prediction_answer": predicted.answer, |
| } |
| for section in SECTIONS: |
| precision, recall, f1 = precision_recall_f1( |
| predicted.canonical_sections[section], |
| expected.canonical_sections[section], |
| ) |
| metrics[f"{section}_precision"] = precision |
| metrics[f"{section}_recall"] = recall |
| metrics[f"{section}_f1"] = f1 |
| all_precision, all_recall, all_f1 = precision_recall_f1(pred_all, exp_all) |
| metrics["all_precision"] = all_precision |
| metrics["all_recall"] = all_recall |
| metrics["all_f1"] = all_f1 |
| return metrics |
|
|
|
|
| def load_jsonl(path: Path) -> list[dict]: |
| rows = [] |
| with path.open() as handle: |
| for line in handle: |
| if line.strip(): |
| rows.append(json.loads(line)) |
| return rows |
|
|
|
|
| def checkpoint_sort_key(path: Path) -> tuple[int, int]: |
| match = STEP_RE.search(path.name) |
| if match: |
| return (1, int(match.group(1))) |
| return (0, -1) |
|
|
|
|
| def checkpoint_label(path: Path) -> str: |
| match = STEP_RE.search(path.name) |
| return f"step {match.group(1)}" if match else "base" |
|
|
|
|
| def mean(values: list[float]) -> float: |
| return sum(values) / len(values) if values else 0.0 |
|
|
|
|
| def summarize(rows: list[dict[str, object]]) -> dict[str, object]: |
| total = len(rows) |
| categories = Counter(str(row["category"]) for row in rows) |
| return { |
| "examples": total, |
| "dominant_exact": sum(bool(row["dominant_exact"]) for row in rows), |
| "dominant_exact_rate": mean([float(bool(row["dominant_exact"])) for row in rows]), |
| "dominant_present_anywhere": sum(bool(row["dominant_present_anywhere"]) for row in rows), |
| "dominant_present_anywhere_rate": mean([float(bool(row["dominant_present_anywhere"])) for row in rows]), |
| "all_exact": sum(bool(row["all_exact"]) for row in rows), |
| "all_exact_rate": mean([float(bool(row["all_exact"])) for row in rows]), |
| "mean_missing_expected": mean([float(row["missing_expected_count"]) for row in rows]), |
| "mean_extra_predicted": mean([float(row["extra_predicted_count"]) for row in rows]), |
| "mean_dominant_f1": mean([float(row["dominant_f1"]) for row in rows]), |
| "mean_irreducible_f1": mean([float(row["irreducible_f1"]) for row in rows]), |
| "mean_reducible_f1": mean([float(row["reducible_f1"]) for row in rows]), |
| "mean_all_f1": mean([float(row["all_f1"]) for row in rows]), |
| "prediction_think_rate": mean([float(bool(row["prediction_has_think"])) for row in rows]), |
| "prediction_answer_tag_rate": mean([float(bool(row["prediction_has_answer"])) for row in rows]), |
| "categories": dict(categories), |
| } |
|
|
|
|
| def annotate_bars(ax: plt.Axes, bars, total: int | None = None) -> None: |
| for bar in bars: |
| height = bar.get_height() |
| label = f"{height:.0f}" |
| if total: |
| label += f"\n{height / total:.0%}" |
| ax.annotate( |
| label, |
| xy=(bar.get_x() + bar.get_width() / 2, height), |
| xytext=(0, 4), |
| textcoords="offset points", |
| ha="center", |
| va="bottom", |
| fontsize=9, |
| ) |
|
|
|
|
| def missing_bucket(value: object) -> str: |
| count = int(value) |
| return "4+" if count >= 4 else str(count) |
|
|
|
|
| def plot_label_text(label: object) -> str: |
| text = str(label) |
| text = text.replace(r"$\\bar{t}$", "tbar").replace(r"$\bar{t}$", "tbar") |
| text = text.replace(r"\\bar{t}", "tbar").replace(r"\bar{t}", "tbar") |
| text = text.replace("$", "").replace("\\", "").replace("{", "").replace("}", "") |
| return re.sub(r"\s+", " ", text).strip() |
|
|
|
|
| def wrapped(labels: list[str], width: int = 18) -> list[str]: |
| wrapped_labels = [] |
| for label in labels: |
| text = plot_label_text(label) |
| wrapped_labels.append("\n".join(textwrap.wrap(text, width=width)) or text) |
| return wrapped_labels |
|
|
|
|
| def plot_latest(rows: list[dict[str, object]], output: Path, title: str, catalog: Catalog) -> None: |
| output.parent.mkdir(parents=True, exist_ok=True) |
| total = len(rows) |
| fig, axes = plt.subplots(2, 2, figsize=(17, 11)) |
| fig.suptitle(title, fontsize=16, y=0.985) |
|
|
| ax = axes[0][0] |
| category_order = [ |
| "dominant exact", |
| "dominant plus extra", |
| "right background, wrong section", |
| "dominant missing", |
| "no dominant parsed", |
| ] |
| counts = Counter(str(row["category"]) for row in rows) |
| values = [counts.get(category, 0) for category in category_order] |
| bars = ax.bar( |
| wrapped(category_order, 14), |
| values, |
| color=["tab:green", "tab:olive", "tab:cyan", "tab:red", "tab:gray"], |
| ) |
| annotate_bars(ax, bars, total) |
| ax.set_title("Dominant Background Outcome") |
| ax.set_ylabel("examples") |
| ax.grid(True, axis="y", alpha=0.25) |
|
|
| ax = axes[0][1] |
| bucket_order = ["0", "1", "2", "3", "4+"] |
| missing_counts = Counter(missing_bucket(row["missing_expected_count"]) for row in rows) |
| bars = ax.bar(bucket_order, [missing_counts.get(bucket, 0) for bucket in bucket_order], color="tab:orange") |
| annotate_bars(ax, bars, total) |
| ax.set_title("Missing Expected Backgrounds") |
| ax.set_xlabel("expected backgrounds absent from prediction") |
| ax.set_ylabel("examples") |
| ax.grid(True, axis="y", alpha=0.25) |
|
|
| ax = axes[1][0] |
| metric_sections = ["dominant", "irreducible", "reducible", "all"] |
| x = list(range(len(metric_sections))) |
| width = 0.24 |
| for offset, metric, color in [ |
| (-width, "precision", "tab:blue"), |
| (0.0, "recall", "tab:purple"), |
| (width, "f1", "tab:green"), |
| ]: |
| values = [mean([float(row[f"{section}_{metric}"]) for row in rows]) for section in metric_sections] |
| ax.bar([idx + offset for idx in x], values, width=width, label=metric, color=color, alpha=0.85) |
| ax.set_xticks(x) |
| ax.set_xticklabels(metric_sections) |
| ax.set_ylim(0, 1.05) |
| ax.set_title("Mean Set-Matching Metrics") |
| ax.set_ylabel("score") |
| ax.grid(True, axis="y", alpha=0.25) |
| ax.legend() |
|
|
| ax = axes[1][1] |
| expected_labels = sorted({str(row["expected_dominant"]) for row in rows}) |
| predicted_labels = sorted({str(row["predicted_dominant"]) or "<none>" for row in rows}) |
| matrix = [] |
| for expected_label in expected_labels: |
| matrix.append( |
| [ |
| sum( |
| 1 |
| for row in rows |
| if str(row["expected_dominant"]) == expected_label |
| and (str(row["predicted_dominant"]) or "<none>") == predicted_label |
| ) |
| for predicted_label in predicted_labels |
| ] |
| ) |
| image = ax.imshow(matrix, cmap="Blues", aspect="auto") |
| ax.set_title("Dominant Confusion Matrix") |
| ax.set_xlabel("predicted dominant") |
| ax.set_ylabel("expected dominant") |
| ax.set_xticks(range(len(predicted_labels))) |
| ax.set_xticklabels(wrapped(predicted_labels, 12), rotation=45, ha="right", fontsize=8) |
| ax.set_yticks(range(len(expected_labels))) |
| ax.set_yticklabels(wrapped(expected_labels, 18), fontsize=8) |
| for y, row_values in enumerate(matrix): |
| for x_idx, value in enumerate(row_values): |
| if value: |
| ax.text(x_idx, y, str(value), ha="center", va="center", fontsize=8) |
| fig.colorbar(image, ax=ax, fraction=0.046, pad=0.04) |
|
|
| summary = summarize(rows) |
| fig.tight_layout(rect=[0, 0.055, 1, 0.955]) |
| fig.text( |
| 0.01, |
| 0.014, |
| ( |
| f"examples {total} | dominant exact {summary['dominant_exact_rate']:.1%} | " |
| f"dominant present anywhere {summary['dominant_present_anywhere_rate']:.1%} | " |
| f"all-background F1 {summary['mean_all_f1']:.3f}" |
| ), |
| ha="left", |
| va="bottom", |
| family="monospace", |
| fontsize=9, |
| ) |
| fig.savefig(output, dpi=180) |
| plt.close(fig) |
|
|
|
|
| def plot_trend(summaries: list[dict[str, object]], output: Path, title: str) -> None: |
| output.parent.mkdir(parents=True, exist_ok=True) |
| labels = [str(row["label"]) for row in summaries] |
| x = list(range(len(labels))) |
|
|
| fig, axes = plt.subplots(2, 2, figsize=(16, 10)) |
| fig.suptitle(title, fontsize=16, y=0.985) |
|
|
| ax = axes[0][0] |
| ax.plot(x, [float(row["dominant_exact_rate"]) for row in summaries], marker="o", label="dominant exact") |
| ax.plot( |
| x, |
| [float(row["dominant_present_anywhere_rate"]) for row in summaries], |
| marker="o", |
| label="dominant present anywhere", |
| ) |
| ax.plot(x, [float(row["all_exact_rate"]) for row in summaries], marker="o", label="all exact") |
| ax.set_ylim(0, 1.05) |
| ax.set_title("Exact Accuracy") |
| ax.set_ylabel("rate") |
| ax.set_xticks(x) |
| ax.set_xticklabels(labels) |
| ax.grid(True, alpha=0.25) |
| ax.legend() |
|
|
| ax = axes[0][1] |
| for key, label in [ |
| ("mean_dominant_f1", "dominant"), |
| ("mean_irreducible_f1", "irreducible"), |
| ("mean_reducible_f1", "reducible"), |
| ("mean_all_f1", "all"), |
| ]: |
| ax.plot(x, [float(row[key]) for row in summaries], marker="o", label=label) |
| ax.set_ylim(0, 1.05) |
| ax.set_title("Mean F1") |
| ax.set_ylabel("F1") |
| ax.set_xticks(x) |
| ax.set_xticklabels(labels) |
| ax.grid(True, alpha=0.25) |
| ax.legend() |
|
|
| ax = axes[1][0] |
| bucket_order = ["0", "1", "2", "3", "4+"] |
| bottoms = [0] * len(labels) |
| colors = ["tab:green", "tab:olive", "tab:orange", "tab:red", "tab:gray"] |
| for bucket, color in zip(bucket_order, colors): |
| values = [int(row.get(f"missing_{bucket}", 0)) for row in summaries] |
| ax.bar(x, values, bottom=bottoms, label=bucket, color=color, alpha=0.85) |
| bottoms = [a + b for a, b in zip(bottoms, values)] |
| ax.set_title("Missing Expected Backgrounds") |
| ax.set_ylabel("examples") |
| ax.set_xticks(x) |
| ax.set_xticklabels(labels) |
| ax.grid(True, axis="y", alpha=0.25) |
| ax.legend(title="missing") |
|
|
| ax = axes[1][1] |
| ax.plot(x, [float(row["mean_missing_expected"]) for row in summaries], marker="o", label="missing") |
| ax.plot(x, [float(row["mean_extra_predicted"]) for row in summaries], marker="o", label="extra") |
| ax.set_title("Mean Set Difference Size") |
| ax.set_ylabel("backgrounds/example") |
| ax.set_xticks(x) |
| ax.set_xticklabels(labels) |
| ax.grid(True, alpha=0.25) |
| ax.legend() |
|
|
| fig.tight_layout(rect=[0, 0.04, 1, 0.955]) |
| fig.savefig(output, dpi=180) |
| plt.close(fig) |
|
|
|
|
| def write_csv(rows: list[dict[str, object]], path: Path) -> None: |
| path.parent.mkdir(parents=True, exist_ok=True) |
| fieldnames = [ |
| "source", |
| "id", |
| "loss", |
| "expected_dominant", |
| "predicted_dominant", |
| "category", |
| "dominant_exact", |
| "dominant_present_anywhere", |
| "all_exact", |
| "missing_expected_count", |
| "extra_predicted_count", |
| "dominant_precision", |
| "dominant_recall", |
| "dominant_f1", |
| "irreducible_precision", |
| "irreducible_recall", |
| "irreducible_f1", |
| "reducible_precision", |
| "reducible_recall", |
| "reducible_f1", |
| "all_precision", |
| "all_recall", |
| "all_f1", |
| "prediction_has_think", |
| "prediction_has_answer", |
| ] |
| with path.open("w", newline="") as handle: |
| writer = csv.DictWriter(handle, fieldnames=fieldnames) |
| writer.writeheader() |
| for row in rows: |
| writer.writerow({field: row.get(field) for field in fieldnames}) |
|
|
|
|
| def write_summary_csv(rows: list[dict[str, object]], path: Path) -> None: |
| path.parent.mkdir(parents=True, exist_ok=True) |
| fieldnames = [ |
| "source", |
| "label", |
| "step", |
| "examples", |
| "dominant_exact", |
| "dominant_exact_rate", |
| "dominant_present_anywhere", |
| "dominant_present_anywhere_rate", |
| "all_exact", |
| "all_exact_rate", |
| "mean_dominant_f1", |
| "mean_irreducible_f1", |
| "mean_reducible_f1", |
| "mean_all_f1", |
| "mean_missing_expected", |
| "mean_extra_predicted", |
| "prediction_think_rate", |
| "prediction_answer_tag_rate", |
| "missing_0", |
| "missing_1", |
| "missing_2", |
| "missing_3", |
| "missing_4+", |
| ] |
| with path.open("w", newline="") as handle: |
| writer = csv.DictWriter(handle, fieldnames=fieldnames) |
| writer.writeheader() |
| for row in rows: |
| writer.writerow({field: row.get(field) for field in fieldnames}) |
|
|
|
|
| def write_traces(rows: list[dict[str, object]], source_rows: list[dict], path: Path) -> None: |
| path.parent.mkdir(parents=True, exist_ok=True) |
| rows_by_id = {str(row["id"]): row for row in rows} |
| with path.open("w") as handle: |
| for source_row in source_rows: |
| row_id = str(source_row.get("id")) |
| metrics = rows_by_id.get(row_id, {}) |
| record = { |
| "id": row_id, |
| "category": metrics.get("category"), |
| "expected_dominant": metrics.get("expected_dominant"), |
| "predicted_dominant": metrics.get("predicted_dominant"), |
| "prediction_think": metrics.get("prediction_think", ""), |
| "prediction_answer": metrics.get("prediction_answer", ""), |
| "prompt": source_row.get("prompt"), |
| "prediction": source_row.get("prediction"), |
| "reference": source_row.get("reference"), |
| } |
| handle.write(json.dumps(record, ensure_ascii=False) + "\n") |
|
|
|
|
| def parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument( |
| "--inputs", |
| nargs="*", |
| type=Path, |
| default=sorted(Path("data/hep_sft/checkpoint_eval").glob("qwen2_5_7b*_val_outputs.jsonl")), |
| help="Generated validation JSONL files. Defaults to qwen2.5 7B checkpoint eval outputs.", |
| ) |
| parser.add_argument("--latest", type=Path, help="Latest checkpoint JSONL. Defaults to highest step among inputs.") |
| parser.add_argument("--catalog", type=Path, default=Path("dataset/config/process_catalog.v1.json")) |
| parser.add_argument("--output-dir", type=Path, default=Path("plotting")) |
| parser.add_argument( |
| "--trend-stem", |
| default="qwen2_5_7b_signature_background_checkpoint_accuracy", |
| help="Filename stem for the across-checkpoint trend plot and CSV.", |
| ) |
| return parser.parse_args() |
|
|
|
|
| def main() -> None: |
| args = parse_args() |
| if not args.inputs: |
| raise SystemExit("No input JSONL files found.") |
| catalog = load_catalog(args.catalog) |
| inputs = sorted(args.inputs, key=checkpoint_sort_key) |
| latest = args.latest or max((path for path in inputs if STEP_RE.search(path.name)), key=checkpoint_sort_key) |
|
|
| all_summaries: list[dict[str, object]] = [] |
| for path in inputs: |
| evaluated = [evaluate_row(row, path, catalog) for row in load_jsonl(path)] |
| summary = summarize(evaluated) |
| missing_counts = Counter(missing_bucket(row["missing_expected_count"]) for row in evaluated) |
| summary.update({f"missing_{bucket}": missing_counts.get(bucket, 0) for bucket in ["0", "1", "2", "3", "4+"]}) |
| summary["source"] = path.name |
| summary["label"] = checkpoint_label(path) |
| sort_group, sort_step = checkpoint_sort_key(path) |
| summary["sort_group"] = sort_group |
| summary["step"] = sort_step if sort_group else None |
| all_summaries.append(summary) |
|
|
| latest_source_rows = load_jsonl(latest) |
| latest_rows = [evaluate_row(row, latest, catalog) for row in latest_source_rows] |
| latest_stem = latest.name.replace("_val_outputs.jsonl", "") |
| latest_plot = args.output_dir / f"{latest_stem}_signature_background_accuracy.png" |
| latest_csv = args.output_dir / f"{latest_stem}_signature_background_examples.csv" |
| latest_summary = args.output_dir / f"{latest_stem}_signature_background_summary.json" |
| latest_traces = args.output_dir / f"{latest_stem}_signature_background_traces.jsonl" |
| trend_plot = args.output_dir / f"{args.trend_stem}.png" |
| trend_csv = args.output_dir / f"{args.trend_stem}.csv" |
|
|
| plot_latest( |
| latest_rows, |
| latest_plot, |
| title=f"{latest_stem} Signature-Background Accuracy", |
| catalog=catalog, |
| ) |
| write_csv(latest_rows, latest_csv) |
| latest_summary.write_text(json.dumps(summarize(latest_rows), indent=2, sort_keys=True) + "\n") |
| write_traces(latest_rows, latest_source_rows, latest_traces) |
|
|
| plot_trend(all_summaries, trend_plot, title="Qwen2.5 7B Signature-Background Accuracy by Checkpoint") |
| write_summary_csv(all_summaries, trend_csv) |
|
|
| print(f"Wrote {latest_plot}") |
| print(f"Wrote {latest_csv}") |
| print(f"Wrote {latest_summary}") |
| print(f"Wrote {latest_traces}") |
| print(f"Wrote {trend_plot}") |
| print(f"Wrote {trend_csv}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|