From 0a3cc65f54d4ffc90b0e73f1f8713820352b27bf Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 28 May 2026 23:11:40 -0500 Subject: Add hidden gradient alignment metric --- scripts/trajectory_mlp_fa.py | 44 ++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 44 insertions(+) (limited to 'scripts/trajectory_mlp_fa.py') diff --git a/scripts/trajectory_mlp_fa.py b/scripts/trajectory_mlp_fa.py index ce20fa5..3b9d5b7 100755 --- a/scripts/trajectory_mlp_fa.py +++ b/scripts/trajectory_mlp_fa.py @@ -44,6 +44,7 @@ class TrajectoryRow: step: int loss: float gradient_cosine: float | None + hidden_gradient_cosine: float | None q_mean: float | None q_min: float | None q_max: float | None @@ -67,6 +68,8 @@ class RunSummary: final_gap_to_bp: float initial_gradient_cosine: float | None final_gradient_cosine: float | None + initial_hidden_gradient_cosine: float | None + final_hidden_gradient_cosine: float | None initial_q_mean: float | None final_q_mean: float | None @@ -309,6 +312,7 @@ def evaluate_bp(weights: list[Array], x: Array, y: Array, step: int) -> Trajecto step=step, loss=loss, gradient_cosine=1.0, + hidden_gradient_cosine=1.0, q_mean=None, q_min=None, q_max=None, @@ -326,6 +330,7 @@ def evaluate_fa( bp_grads, loss = gradients(weights, x, y, feedback=None) fa_grads, _ = gradients(weights, x, y, feedback=feedback) grad_cos = cosine(flatten(bp_grads), flatten(fa_grads)) + hidden_grad_cos = cosine(flatten(bp_grads[:-1]), flatten(fa_grads[:-1])) layer_cosines = layer_gradient_cosines(bp_grads, fa_grads) q_values = layer_q_alignments(weights, feedback) @@ -335,6 +340,7 @@ def evaluate_fa( step=step, loss=loss, gradient_cosine=grad_cos, + hidden_gradient_cosine=hidden_grad_cos, q_mean=float(np.mean(q_values)), q_min=float(np.min(q_values)), q_max=float(np.max(q_values)), @@ -426,6 +432,8 @@ def make_summaries( final_gap_to_bp=0.0, initial_gradient_cosine=1.0, final_gradient_cosine=1.0, + initial_hidden_gradient_cosine=1.0, + final_hidden_gradient_cosine=1.0, initial_q_mean=None, final_q_mean=None, ) @@ -442,6 +450,8 @@ def make_summaries( final_gap_to_bp=final.loss - bp_final, initial_gradient_cosine=first.gradient_cosine, final_gradient_cosine=final.gradient_cosine, + initial_hidden_gradient_cosine=first.hidden_gradient_cosine, + final_hidden_gradient_cosine=final.hidden_gradient_cosine, initial_q_mean=first.q_mean, final_q_mean=final.q_mean, ) @@ -526,6 +536,29 @@ def save_plots(trajectories: list[TrajectoryRow], outdir: Path) -> list[Path]: plt.close() paths.append(gamma_path) + hidden_gamma_path = outdir / "hidden_gradient_cosine_curves.png" + plt.figure(figsize=(7, 4.5)) + for seed in fa_seeds: + rows = [ + row + for row in trajectories + if row.run_type == "fa" and row.feedback_seed == seed + ] + plt.plot( + [row.step for row in rows], + [row.hidden_gradient_cosine for row in rows], + alpha=0.75, + label=f"seed={seed}", + ) + plt.axhline(0.0, color="black", linewidth=1) + plt.xlabel("step") + plt.ylabel("cos(BP hidden gradient, FA hidden gradient)") + plt.title("Hidden-layer surrogate gradient alignment") + plt.tight_layout() + plt.savefig(hidden_gamma_path, dpi=180) + plt.close() + paths.append(hidden_gamma_path) + q_path = outdir / "q_alignment_curves.png" plt.figure(figsize=(7, 4.5)) for seed in fa_seeds: @@ -588,6 +621,11 @@ def main() -> None: fa_final_gammas = [ row.final_gradient_cosine for row in summaries if row.run_type == "fa" ] + fa_final_hidden_gammas = [ + row.final_hidden_gradient_cosine + for row in summaries + if row.run_type == "fa" + ] print(f"widths: {widths}") print(f"bp_final_loss: {bp_final:.8g}") print( @@ -602,6 +640,12 @@ def main() -> None: f"min={np.min(fa_final_gammas):.8g}, " f"max={np.max(fa_final_gammas):.8g}" ) + print( + "fa_final_hidden_gradient_cosine: " + f"mean={np.mean(fa_final_hidden_gammas):.8g}, " + f"min={np.min(fa_final_hidden_gammas):.8g}, " + f"max={np.max(fa_final_hidden_gammas):.8g}" + ) print(f"summary: {outdir / 'summary.csv'}") print(f"trajectories: {outdir / 'trajectories.csv'}") print(f"layer_metrics: {outdir / 'layer_metrics.csv'}") -- cgit v1.2.3