hep-posttraining / plotting /plot_sft_eval_comparison.py
ho22joshua's picture
Upload folder using huggingface_hub
587d4ca verified
Raw
History Blame Contribute Delete
8.67 kB
#!/usr/bin/env python3
"""Compare base-model and SFT-adapter validation results."""
from __future__ import annotations
import argparse
import csv
import json
import math
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
def load_json(path: Path) -> dict:
return json.loads(path.read_text())
def load_jsonl(path: Path) -> list[dict]:
rows = []
with path.open() as handle:
for line in handle:
line = line.strip()
if line:
rows.append(json.loads(line))
return rows
def word_count(value: object) -> int:
return len(str(value or "").split())
def by_scored_id(path: Path) -> dict[str, dict]:
return {str(row["id"]): row for row in load_jsonl(path) if row.get("status") == "scored" and "id" in row}
def generations_by_id(path: Path) -> dict[str, dict]:
return {str(row["id"]): row for row in load_jsonl(path) if "id" in row}
def build_rows(
base_examples: Path,
sft_examples: Path,
base_generations: Path | None,
sft_generations: Path | None,
) -> list[dict]:
base_loss = by_scored_id(base_examples)
sft_loss = by_scored_id(sft_examples)
shared_ids = sorted(set(base_loss) & set(sft_loss))
if not shared_ids:
raise SystemExit("No shared scored example ids found between base and SFT example files.")
base_gen = generations_by_id(base_generations) if base_generations else {}
sft_gen = generations_by_id(sft_generations) if sft_generations else {}
rows = []
for example_id in shared_ids:
base_row = base_loss[example_id]
sft_row = sft_loss[example_id]
base_loss_value = float(base_row["loss"])
sft_loss_value = float(sft_row["loss"])
base_generation = base_gen.get(example_id, {})
sft_generation = sft_gen.get(example_id, {})
reference = sft_generation.get("reference") or base_generation.get("reference")
rows.append(
{
"id": example_id,
"base_loss": base_loss_value,
"sft_loss": sft_loss_value,
"loss_delta": sft_loss_value - base_loss_value,
"loss_reduction_pct": (
100.0 * (base_loss_value - sft_loss_value) / base_loss_value
if base_loss_value
else math.nan
),
"base_assistant_tokens": int(base_row.get("assistant_tokens", 0)),
"sft_assistant_tokens": int(sft_row.get("assistant_tokens", 0)),
"base_prediction_words": word_count(base_generation.get("prediction")),
"sft_prediction_words": word_count(sft_generation.get("prediction")),
"reference_words": word_count(reference),
}
)
return sorted(rows, key=lambda row: row["base_loss"], reverse=True)
def mean(values: list[float]) -> float | None:
return sum(values) / len(values) if values else None
def write_csv(rows: list[dict], path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
fieldnames = [
"id",
"base_loss",
"sft_loss",
"loss_delta",
"loss_reduction_pct",
"base_assistant_tokens",
"sft_assistant_tokens",
"base_prediction_words",
"sft_prediction_words",
"reference_words",
]
with path.open("w", newline="") as handle:
writer = csv.DictWriter(handle, fieldnames=fieldnames)
writer.writeheader()
writer.writerows(rows)
def annotate_bars(ax: plt.Axes, bars: list, fmt: str = "{:.3g}") -> None:
for bar in bars:
height = bar.get_height()
ax.annotate(
fmt.format(height),
xy=(bar.get_x() + bar.get_width() / 2, height),
xytext=(0, 4),
textcoords="offset points",
ha="center",
va="bottom",
fontsize=9,
)
def plot(
base_summary: dict,
sft_summary: dict,
rows: list[dict],
output: Path,
title: str,
base_label: str,
sft_label: str,
) -> None:
output.parent.mkdir(parents=True, exist_ok=True)
fig, axes = plt.subplots(2, 2, figsize=(15, 10))
fig.suptitle(title, fontsize=16, y=0.985)
ax = axes[0][0]
loss_bars = ax.bar(
[base_label, sft_label],
[float(base_summary["loss"]), float(sft_summary["loss"])],
color=["tab:gray", "tab:blue"],
)
annotate_bars(ax, loss_bars)
ax.set_title("Validation Loss")
ax.set_ylabel("assistant-token NLL")
ax.grid(True, axis="y", alpha=0.25)
ax.text(
0.5,
0.92,
(
f"perplexity: {float(base_summary['perplexity']):.3g} -> "
f"{float(sft_summary['perplexity']):.3g}"
),
transform=ax.transAxes,
ha="center",
va="top",
fontsize=10,
)
ax = axes[0][1]
indices = list(range(1, len(rows) + 1))
base_losses = [row["base_loss"] for row in rows]
sft_losses = [row["sft_loss"] for row in rows]
for idx, base_value, sft_value in zip(indices, base_losses, sft_losses):
ax.plot([idx, idx], [base_value, sft_value], color="0.82", linewidth=1)
ax.scatter(indices, base_losses, s=28, label=base_label, color="tab:gray")
ax.scatter(indices, sft_losses, s=28, label=sft_label, color="tab:blue")
ax.set_title("Per-Example Loss")
ax.set_xlabel("examples sorted by base loss")
ax.set_ylabel("loss")
ax.grid(True, alpha=0.25)
ax.legend()
ax = axes[1][0]
reductions = [row["loss_reduction_pct"] for row in rows if math.isfinite(row["loss_reduction_pct"])]
ax.hist(reductions, bins=min(12, max(4, len(reductions) // 2)), color="tab:green", alpha=0.8)
ax.axvline(mean(reductions), color="black", linestyle="--", linewidth=1.5, label="mean")
ax.set_title("Per-Example Loss Reduction")
ax.set_xlabel("reduction vs base (%)")
ax.set_ylabel("examples")
ax.grid(True, axis="y", alpha=0.25)
ax.legend()
ax = axes[1][1]
length_labels = [base_label, sft_label, "Reference"]
length_values = [
mean([row["base_prediction_words"] for row in rows]) or 0.0,
mean([row["sft_prediction_words"] for row in rows]) or 0.0,
mean([row["reference_words"] for row in rows]) or 0.0,
]
length_bars = ax.bar(length_labels, length_values, color=["tab:gray", "tab:blue", "tab:orange"])
annotate_bars(ax, length_bars, fmt="{:.1f}")
ax.set_title("Generated Answer Length")
ax.set_ylabel("mean words")
ax.grid(True, axis="y", alpha=0.25)
loss_reduction = 100.0 * (
float(base_summary["loss"]) - float(sft_summary["loss"])
) / float(base_summary["loss"])
summary = (
f"examples {len(rows)} | "
f"loss {float(base_summary['loss']):.4f} -> {float(sft_summary['loss']):.4f} "
f"({loss_reduction:.1f}% lower) | "
f"ppl {float(base_summary['perplexity']):.4f} -> {float(sft_summary['perplexity']):.4f}"
)
fig.tight_layout(rect=[0, 0.045, 1, 0.955])
fig.text(0.01, 0.012, summary, ha="left", va="bottom", family="monospace", fontsize=9)
fig.savefig(output, dpi=180)
plt.close(fig)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Plot base vs SFT validation-result comparisons.")
parser.add_argument("--base-summary", required=True, type=Path)
parser.add_argument("--sft-summary", required=True, type=Path)
parser.add_argument("--base-examples", required=True, type=Path)
parser.add_argument("--sft-examples", required=True, type=Path)
parser.add_argument("--base-generations", type=Path)
parser.add_argument("--sft-generations", type=Path)
parser.add_argument("--output", required=True, type=Path)
parser.add_argument("--csv", type=Path)
parser.add_argument("--title", default="Qwen2.5 7B Base vs LoRA-16 SFT")
parser.add_argument("--base-label", default="Base")
parser.add_argument("--sft-label", default="LoRA-16 SFT")
return parser.parse_args()
def main() -> None:
args = parse_args()
base_summary = load_json(args.base_summary)
sft_summary = load_json(args.sft_summary)
rows = build_rows(
args.base_examples,
args.sft_examples,
args.base_generations,
args.sft_generations,
)
plot(
base_summary,
sft_summary,
rows,
args.output,
args.title,
args.base_label,
args.sft_label,
)
print(f"Wrote {args.output}")
if args.csv:
write_csv(rows, args.csv)
print(f"Wrote {args.csv}")
if __name__ == "__main__":
main()