"""Multi-seed LLM robustness figure — the recombination claims with error bars. Aggregates the three multi-seed 0.5B experiments (fixed test sets, training seed varied) into one figure with 95% CIs over seeds: (A) llm_merge_seeds — Fisher–Muller: merged specialists vs the best single specialist, overall and worst-family (5 seeds). (B) llm_moe_hard_seeds — union (routing) vs fusion (soup/ties) on the hard benchmark (3 seeds). (C) llm_directed_hard_seeds — directed offspring selection vs the a-priori soup, hard (3 seeds). Usage: python figures/plot_llm_seeds.py # the 0.5B seed bundles (default) python figures/plot_llm_seeds.py --merge results/llm_merge_hpc --moe results/llm_moe_hard_hpc \ --directed results/llm_directed_hard_hpc --out results/llm_merge_hpc --tag 7B Bundles may be flat (one seed) or per-seed sub-bundles ``s{seed}/`` (HPC array layout). """ from __future__ import annotations import sys from pathlib import Path import numpy as np import matplotlib.pyplot as plt sys.path.insert(0, str(Path(__file__).parent)) from _figlib import load_seed_bundles, savefig # noqa: E402 def _agg(df, models, metric): """Per-model mean and 95% CI over seeds for one metric.""" out = [] for m in models: v = df[(df["model"] == m) & (df["metric"] == metric)].groupby("seed")["accuracy"].mean() out.append((v.mean(), 1.96 * v.std(ddof=1) / np.sqrt(len(v)) if len(v) > 1 else 0.0)) return out def _panel(ax, df, models, labels, title): x = np.arange(len(models)) for off, metric, color in ((-0.17, "overall", "#2c7fb8"), (0.17, "worst_family", "#d62728")): vals = _agg(df, models, metric) ax.bar(x + off, [v for v, _ in vals], 0.34, yerr=[e for _, e in vals], capsize=3, color=color, label=metric) ax.set_xticks(x); ax.set_xticklabels(labels, fontsize=8) ax.set(ylabel="verifier accuracy", ylim=(0, 1.0), title=title) ax.legend(frameon=False, fontsize=8) def _best_spec(df): """Synthesise a best-single-specialist row set per seed (max over spec_* by overall).""" specs = sorted(m for m in df["model"].unique() if m.startswith("spec_")) rows = [] for s, sub in df.groupby("seed"): ov = {m: sub[(sub["model"] == m) & (sub["metric"] == "overall")]["accuracy"].mean() for m in specs} best = max(ov, key=ov.get) b = sub[sub["model"] == best].copy() b["model"] = "best_specialist" rows.append(b) import pandas as pd return pd.concat([df] + rows, ignore_index=True) def main(merge="results/llm_merge_seeds", moe="results/llm_moe_hard_seeds", directed="results/llm_directed_hard_seeds", out="results/llm_merge_seeds", tag="0.5B") -> None: fig, axes = plt.subplots(1, 3, figsize=(16, 4.8)) df, _ = load_seed_bundles(merge) n = df["seed"].nunique() _panel(axes[0], _best_spec(df), ["base", "best_specialist", "merge_soup", "merge_ties"], ["base", "best\nspecialist", "merge\n(soup)", "merge\n(ties)"], f"(A) Fisher–Muller with error bars\n({n} seeds, easy benchmark, {tag})") df, _ = load_seed_bundles(moe) n = df["seed"].nunique() _panel(axes[1], _best_spec(df), ["best_specialist", "merge_soup", "merge_ties", "moe_oracle", "moe_learned"], ["best\nspecialist", "fusion\n(soup)", "fusion\n(ties)", "union\n(route,oracle)", "union\n(route,learned)"], f"(B) union vs fusion, hard benchmark\n({n} seeds, {tag})") df, _ = load_seed_bundles(directed) n = df["seed"].nunique() _panel(axes[2], df, ["merge_soup", "directed_overall", "directed_balanced"], ["a-priori soup", "directed\n(overall)", "directed\n(balanced)"], f"(C) directed offspring selection, hard\n({n} seeds, {tag})") fig.suptitle("The LLM recombination claims are seed-robust (fixed test sets; training seed varied; 95% CI)", y=1.03, fontsize=12) fig.tight_layout() savefig(fig, out, "llm_seeds") if __name__ == "__main__": import argparse ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("--merge", default="results/llm_merge_seeds") ap.add_argument("--moe", default="results/llm_moe_hard_seeds") ap.add_argument("--directed", default="results/llm_directed_hard_seeds") ap.add_argument("--out", default="results/llm_merge_seeds") ap.add_argument("--tag", default="0.5B") main(**vars(ap.parse_args()))