summaryrefslogtreecommitdiff
path: root/scripts/trajectory_mlp_fa.py
diff options
context:
space:
mode:
Diffstat (limited to 'scripts/trajectory_mlp_fa.py')
-rwxr-xr-xscripts/trajectory_mlp_fa.py44
1 files changed, 44 insertions, 0 deletions
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'}")