summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-06-02 13:36:20 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-06-02 13:36:20 -0500
commiteea755edb954c3fe154b354dec7870fb529f498f (patch)
treee40d4440fcf6e54c7bcab089d91175b25207d895
parente666b35931f072dc6b0ece5823c689db5cf0a0c9 (diff)
Expand empirical transition overlays
-rw-r--r--scripts/downstream_capacity_sweep.py27
-rw-r--r--scripts/plot_capacity_transition_overlay.py109
2 files changed, 78 insertions, 58 deletions
diff --git a/scripts/downstream_capacity_sweep.py b/scripts/downstream_capacity_sweep.py
index 6e42550..a5915ea 100644
--- a/scripts/downstream_capacity_sweep.py
+++ b/scripts/downstream_capacity_sweep.py
@@ -37,11 +37,14 @@ class RunConfig:
optimizer: str
init_seeds: int
feedback_seeds: int
+ init_seed_offset: int
+ feedback_seed_offset: int
data_seed: int
noise_std: float
feedback_scale: str
capacity_q: float
jacobian_lambda_rel: float
+ skip_jacobian: bool
device: str
torch_threads: int
outdir: str
@@ -120,6 +123,8 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--optimizer", choices=["sgd", "adam"], default="sgd")
parser.add_argument("--init-seeds", type=int, default=3)
parser.add_argument("--feedback-seeds", type=int, default=5)
+ parser.add_argument("--init-seed-offset", type=int, default=0)
+ parser.add_argument("--feedback-seed-offset", type=int, default=0)
parser.add_argument("--data-seed", type=int, default=0)
parser.add_argument("--noise-std", type=float, default=0.0)
parser.add_argument(
@@ -129,6 +134,7 @@ def parse_args() -> argparse.Namespace:
)
parser.add_argument("--capacity-q", type=float, default=0.01)
parser.add_argument("--jacobian-lambda-rel", type=float, default=1e-3)
+ parser.add_argument("--skip-jacobian", action="store_true")
parser.add_argument("--device", choices=["cpu", "cuda"], default="cpu")
parser.add_argument("--torch-threads", type=int, default=0)
parser.add_argument(
@@ -158,11 +164,14 @@ def parse_config(args: argparse.Namespace) -> RunConfig:
optimizer=args.optimizer,
init_seeds=args.init_seeds,
feedback_seeds=args.feedback_seeds,
+ init_seed_offset=args.init_seed_offset,
+ feedback_seed_offset=args.feedback_seed_offset,
data_seed=args.data_seed,
noise_std=args.noise_std,
feedback_scale=args.feedback_scale,
capacity_q=args.capacity_q,
jacobian_lambda_rel=args.jacobian_lambda_rel,
+ skip_jacobian=args.skip_jacobian,
device=args.device,
torch_threads=args.torch_threads,
outdir=str(args.outdir),
@@ -792,7 +801,7 @@ def main() -> None:
flush=True,
)
for init_index in range(config.init_seeds):
- init_seed = 10_000 + init_index
+ init_seed = 10_000 + config.init_seed_offset + init_index
initial_weights = initialize_weights(config, width, init_seed)
bp_weights = train(
initial_weights,
@@ -805,9 +814,12 @@ def main() -> None:
)
bp_train = mse(bp_weights, x_train, y_train)
bp_test = mse(bp_weights, x_test, y_test)
- d_eff, hard_rank, lam = jacobian_effective_dimension(
- bp_weights, x_probe, config.jacobian_lambda_rel
- )
+ if config.skip_jacobian:
+ d_eff, hard_rank, lam = 0.0, 0, 0.0
+ else:
+ d_eff, hard_rank, lam = jacobian_effective_dimension(
+ bp_weights, x_probe, config.jacobian_lambda_rel
+ )
redundancy = p_count - d_eff - burden
rows.append(
RunRow(
@@ -835,7 +847,12 @@ def main() -> None:
)
for feedback_index in range(config.feedback_seeds):
- feedback_seed = 100_000 + init_index * 1000 + feedback_index
+ feedback_seed = (
+ 100_000
+ + config.feedback_seed_offset
+ + init_index * 1000
+ + feedback_index
+ )
feedback = init_feedback(config, width, feedback_seed)
fa_weights = train(
initial_weights,
diff --git a/scripts/plot_capacity_transition_overlay.py b/scripts/plot_capacity_transition_overlay.py
index e636725..cc53012 100644
--- a/scripts/plot_capacity_transition_overlay.py
+++ b/scripts/plot_capacity_transition_overlay.py
@@ -28,6 +28,7 @@ class SummaryRow:
@dataclass(frozen=True)
class RunRow:
+ source_id: int
width: float
feedback_seed: int
run_type: str
@@ -42,6 +43,7 @@ def parse_args() -> argparse.Namespace:
type=Path,
default=Path("outputs/downstream_capacity_random_main_fast"),
)
+ parser.add_argument("--extra-run-dir", type=Path, action="append", default=[])
parser.add_argument("--samples", type=int, default=2048)
parser.add_argument(
"--transition-width",
@@ -73,12 +75,13 @@ def read_summary(path: Path) -> list[SummaryRow]:
return rows
-def read_runs(path: Path) -> list[RunRow]:
+def read_runs(path: Path, source_id: int) -> list[RunRow]:
rows: list[RunRow] = []
with path.open() as handle:
for row in csv.DictReader(handle):
rows.append(
RunRow(
+ source_id=source_id,
width=float(row["width"]),
feedback_seed=int(row["feedback_seed"]),
run_type=row["run_type"],
@@ -161,6 +164,7 @@ def calibrate_scale(summary: list[SummaryRow], raw_mean_at_widths: np.ndarray) -
def plot_transition(
run_dir: Path,
+ extra_run_dirs: list[Path],
samples: int,
transition_width: float,
outname: str,
@@ -168,7 +172,9 @@ def plot_transition(
xmax: float,
) -> Path:
summary = read_summary(run_dir / "width_summary.csv")
- runs = read_runs(run_dir / "runs.csv")
+ runs = read_runs(run_dir / "runs.csv", 0)
+ for source_id, extra_run_dir in enumerate(extra_run_dirs, start=1):
+ runs.extend(read_runs(extra_run_dir / "runs.csv", source_id))
input_dim, output_dim, task_dim = infer_dims(run_dir)
rng = np.random.default_rng(0)
@@ -193,71 +199,67 @@ def plot_transition(
dense_lo = scale * dense_lo_raw
dense_hi = scale * dense_hi_raw
- summary_margin = np.array([row.fa_capacity_margin for row in summary], dtype=np.float64)
- summary_gap = np.array([row.train_gap_mean for row in summary], dtype=np.float64)
-
- by_seed: dict[int, list[RunRow]] = defaultdict(list)
+ by_seed: dict[tuple[int, int], list[RunRow]] = defaultdict(list)
for row in runs:
if row.run_type == "fa":
- by_seed[row.feedback_seed].append(row)
+ by_seed[(row.source_id, row.feedback_seed)].append(row)
fig, axes = plt.subplots(1, 2, figsize=(13.5, 5.2), sharex=True, sharey=True)
- for axis, show_trajectories in zip(axes, [False, True]):
- axis.fill_between(
- dense_margin,
- dense_lo,
- dense_hi,
- color="tab:orange",
- alpha=0.22,
- linewidth=0,
- label="theory 10-90% band",
- )
- axis.plot(
- dense_margin,
- dense_mean,
- color="black",
- linewidth=2.4,
- label="theory mean",
- )
- axis.axvline(0.0, color="black", linestyle="--", linewidth=1.1, alpha=0.8)
+ axes[0].fill_between(
+ dense_margin,
+ dense_lo,
+ dense_hi,
+ color="tab:orange",
+ alpha=0.22,
+ linewidth=0,
+ label="theory 10-90% band",
+ )
+ axes[0].plot(
+ dense_margin,
+ dense_mean,
+ color="black",
+ linewidth=2.4,
+ label="theory mean",
+ )
+ axes[0].axvline(0.0, color="black", linestyle="--", linewidth=1.1, alpha=0.8)
+
+ for axis in axes:
axis.axhline(0.0, color="black", linewidth=0.9, alpha=0.75)
axis.set_xlim(xmin, xmax)
axis.set_xlabel("hard FA capacity margin P - K_FA - N*out")
axis.grid(True, color="#dddddd", linewidth=0.8, alpha=0.7)
- if show_trajectories:
- for seed, rows in sorted(by_seed.items()):
- ordered = sorted(rows, key=lambda item: item.fa_capacity_margin)
- axis.plot(
- [row.fa_capacity_margin for row in ordered],
- [row.train_gap_to_bp for row in ordered],
- color="tab:blue",
- alpha=0.12,
- linewidth=0.9,
- )
- axis.plot(
- summary_margin,
- summary_gap,
- color="tab:blue",
- linewidth=2.2,
- alpha=0.95,
- label="empirical mean",
- )
- axis.scatter(
- summary_margin,
- summary_gap,
- color="black",
- s=34,
- zorder=4,
- label="empirical points",
- )
+ for _seed, rows in sorted(by_seed.items()):
+ ordered = sorted(rows, key=lambda item: item.fa_capacity_margin)
+ axes[1].plot(
+ [row.fa_capacity_margin for row in ordered],
+ [row.train_gap_to_bp for row in ordered],
+ color="tab:blue",
+ alpha=0.07,
+ linewidth=0.85,
+ )
+
+ axes[1].text(
+ 0.98,
+ 0.96,
+ f"{len(by_seed)} empirical FA trajectories",
+ transform=axes[1].transAxes,
+ ha="right",
+ va="top",
+ fontsize=10,
+ bbox={
+ "boxstyle": "round,pad=0.25",
+ "facecolor": "white",
+ "edgecolor": "#cccccc",
+ "alpha": 0.85,
+ },
+ )
axes[0].set_title("P1: smooth capacity-transition prediction")
- axes[1].set_title("P2: transparent trajectory ensemble overlay")
+ axes[1].set_title("P2: empirical trajectory overlay")
axes[0].set_ylabel("FA train MSE - BP train MSE")
axes[0].legend(loc="upper right")
- axes[1].legend(loc="upper right")
fig.suptitle(
f"Random-label memorization transition, scale={scale:.4g}, transition width={transition_width:g}",
y=1.02,
@@ -274,6 +276,7 @@ def main() -> None:
args = parse_args()
outpath = plot_transition(
args.run_dir,
+ args.extra_run_dir,
args.samples,
args.transition_width,
args.outname,