From 2f3f44d49df3d1dd42200bae5fa6cbe943c3d42c Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Tue, 21 Jul 2026 07:26:54 -0500 Subject: feat: add magnitude-matched raw apical control --- experiments/nuisance_control_sweep.sh | 66 +++++++++++++++++++++++++++++++++++ experiments/run.py | 5 ++- experiments/smoke.py | 22 +++++++++++- 3 files changed, 91 insertions(+), 2 deletions(-) create mode 100755 experiments/nuisance_control_sweep.sh (limited to 'experiments') diff --git a/experiments/nuisance_control_sweep.sh b/experiments/nuisance_control_sweep.sh new file mode 100755 index 0000000..ef715f4 --- /dev/null +++ b/experiments/nuisance_control_sweep.sh @@ -0,0 +1,66 @@ +#!/usr/bin/env bash +# Three-way residualization control with per-neuron soma-dendrite coupling: +# raw = unmodified apical teaching signal +# matched = raw direction, per-sample norm matched to the innovation +# residual = soma-dendritic innovation +# +# Usage: +# bash experiments/nuisance_control_sweep.sh \ +# "" "" [pert_mode] [n_dirs] +# +# Example: +# bash experiments/nuisance_control_sweep.sh \ +# 5 "0 0.05 0.2 0.5" "0 1 2" 15 nuis_ctrl simultaneous 16 +set -eu + +cd "$(dirname "$0")/.." +PY=/home/yurenh2/miniconda3/envs/ep_pascal/bin/python3 +GPU="${1:?GPU index required}" +RHOS="${2:-0 0.05 0.2 0.5}" +SEEDS="${3:-0}" +EPOCHS="${4:-15}" +PREFIX="${5:-nuis_ctrl}" +PERT_MODE="${6:-simultaneous}" +N_DIRS="${7:-16}" + +export CUDA_VISIBLE_DEVICES="$GPU" +export OMP_NUM_THREADS=2 +mkdir -p results logs/nuisance + +run_one() { + local rho="$1" + local arm="$2" + local seed="$3" + local residual=0 + local scale_control=none + local rho_tag="${rho//./p}" + local tag="${PREFIX}_rho${rho_tag}_${arm}_s${seed}" + local result="results/${tag}.json" + local log="logs/nuisance/${tag}.log" + if [[ "$arm" == residual ]]; then + residual=1 + elif [[ "$arm" == matched ]]; then + scale_control=match_innovation_norm + fi + if [[ -f "$result" ]]; then + echo "skip $tag (result exists)" + return + fi + echo ">>> $tag $(date --iso-8601=seconds) gpu=$GPU" + "$PY" experiments/run.py \ + --mode sdil --dataset mnist --depth 3 --width 256 --epochs "$EPOCHS" \ + --seed "$seed" --nuis_rho "$rho" --use_residual "$residual" \ + --raw_scale_control "$scale_control" --predictor_mode diagonal \ + --p_warmup_steps 200 --p_warmup_eta 0.05 \ + --pert_mode "$PERT_MODE" --pert_ndirs "$N_DIRS" --log_every 100 \ + --tag "$tag" --outdir results > "$log" 2>&1 + grep -h DONE "$log" +} + +for seed in $SEEDS; do + for rho in $RHOS; do + for arm in raw matched residual; do + run_one "$rho" "$arm" "$seed" + done + done +done diff --git a/experiments/run.py b/experiments/run.py index 3cb2e00..eea3f91 100644 --- a/experiments/run.py +++ b/experiments/run.py @@ -65,7 +65,8 @@ def build(args, device): momentum=args.momentum, settle_steps=args.settle_steps, kappa=args.kappa, feedback=args.feedback, p_update_on_neutral=bool(args.p_neutral), - normalize_delta=bool(args.normalize_delta)) + normalize_delta=bool(args.normalize_delta), + raw_scale_control=args.raw_scale_control) else: raise ValueError(args.mode) return net, cfg @@ -174,6 +175,8 @@ def get_args(): p.add_argument("--pert_ndirs", type=int, default=4) p.add_argument("--pert_mode", default="layerwise", choices=["layerwise", "simultaneous"]) p.add_argument("--use_residual", type=int, default=1) + p.add_argument("--raw_scale_control", default="none", + choices=["none", "match_innovation_norm"]) p.add_argument("--learn_A", type=int, default=1) p.add_argument("--learn_P", type=int, default=1) p.add_argument("--p_neutral", type=int, default=1) # P update on neutral (c=0) drive diff --git a/experiments/smoke.py b/experiments/smoke.py index 575f0bd..0bdb0ca 100644 --- a/experiments/smoke.py +++ b/experiments/smoke.py @@ -8,7 +8,7 @@ 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, simultaneous_node_perturbation_targets, - neutral_p_update) + neutral_p_update, teaching_signal) from sdil.baselines import dfa_config from sdil import probes from sdil.data import get_dataset, onehot @@ -164,6 +164,26 @@ def main(): assert al_raw["cos_r_negg"] == al_raw["cos_apical_negg"], ( "reported teaching alignment must honor use_residual=False") + # A magnitude-matched raw signal is still pointed along raw apical activity, + # but has exactly the innovation's norm for every sample. This isolates the + # directional effect of subtracting the soma-predictable component. + cfg_matched = SDILConfig(use_residual=False, learn_A=False, learn_P=True, + raw_scale_control="match_innovation_norm") + with torch.no_grad(): + fwd_n = net_n.forward(x) + error_n = net_n.output_error(fwd_n["h"][-1], yoh) + for l in range(net_n.L - 1): + matched, raw, innovation = teaching_signal( + net_n, l, error_n, fwd_n["h"][l + 1], cfg_matched) + assert torch.allclose(matched.norm(dim=1), innovation.norm(dim=1), + atol=1e-6, rtol=1e-5) + assert row_cos(matched, raw) > 0.99999 + al_matched = probes.alignment_report(net_n, x, y, yoh, cfg_matched) + for matched_cos, raw_cos in zip(al_matched["cos_r_negg"], + al_matched["cos_apical_negg"]): + assert abs(matched_cos - raw_cos) < 1e-6 + print(" matched raw: direction preserved; per-sample innovation norm matched") + # ---- CHECK 5: single-step loss-decrease ratio vs exact GD ---- print("\nCHECK5 single-step loss-decrease ratio vs exact GD:") ldr = probes.loss_decrease_ratio(net, x, y, yoh, cfg, step=100) -- cgit v1.2.3