From 0f107d1ee4e750421a0e95f8bd06bf0f629d2fa7 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Tue, 21 Jul 2026 07:02:04 -0500 Subject: feat: multiplex causal probes across hidden layers --- experiments/smoke.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) (limited to 'experiments/smoke.py') diff --git a/experiments/smoke.py b/experiments/smoke.py index 8196d4f..f440f7f 100644 --- a/experiments/smoke.py +++ b/experiments/smoke.py @@ -7,7 +7,8 @@ import torch sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from sdil.core import (SDILNet, SDILConfig, sdil_step, - node_perturbation_targets, neutral_p_update) + node_perturbation_targets, simultaneous_node_perturbation_targets, + neutral_p_update) from sdil.baselines import dfa_config from sdil import probes from sdil.data import get_dataset, onehot @@ -104,6 +105,11 @@ def main(): print(f"CHECK1 node-pert n_dirs={ndirs:2d} cos(q,-g) per layer = " + " ".join(f"{c:+.3f}" for c in cs)) assert all(c > 0 for c in cs), "node perturbation q should align with -grad" + qs = simultaneous_node_perturbation_targets(net, x, y, sigma=1e-2, n_dirs=32) + cs = [row_cos(qs[l], -grads[l]) for l in range(len(qs))] + print("CHECK1 simultaneous n_dirs=32 cos(q,-g) per layer = " + + " ".join(f"{c:+.3f}" for c in cs)) + assert all(c > 0.2 for c in cs), "simultaneous perturbation should align in expectation" # ---- CHECK 2: SDIL overfits a fixed batch and alignment climbs ---- net = SDILNet(sizes, device=dev, seed=2) -- cgit v1.2.3