summaryrefslogtreecommitdiff
path: root/ep_run/casc_bp_train.py
diff options
context:
space:
mode:
authorYuren Hao <yurenh2@illinois.edu>2026-07-09 09:52:33 -0500
committerYuren Hao <yurenh2@illinois.edu>2026-07-09 09:52:33 -0500
commit32d1fdbfd0becef18351762d5d1d2924640a52be (patch)
tree671841b24d3284cd4018f64903583d170dd289c2 /ep_run/casc_bp_train.py
parent7e338ed314d7d73fdfaf32daa5bc105127a42c99 (diff)
cascade root cause found: tied readout + default N(0,1) embedding init = pathological top-CE stiffness once predictions sharpen (sigma_tok~76 vs GPT-standard 0.02 giving ~1.6); add --untie and --tok_init; diag arms A/B/C
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn
Diffstat (limited to 'ep_run/casc_bp_train.py')
-rw-r--r--ep_run/casc_bp_train.py4
1 files changed, 4 insertions, 0 deletions
diff --git a/ep_run/casc_bp_train.py b/ep_run/casc_bp_train.py
index 8e039ec..5194c45 100644
--- a/ep_run/casc_bp_train.py
+++ b/ep_run/casc_bp_train.py
@@ -14,6 +14,7 @@ ap.add_argument('--lr', type=float, default=3e-4); ap.add_argument('--warmup', t
ap.add_argument('--seed', type=int, default=0)
ap.add_argument('--save_every', type=int, default=500); ap.add_argument('--log', type=int, default=200)
ap.add_argument('--wandb', default=''); ap.add_argument('--wandb_run', default='')
+ap.add_argument('--tok_init', type=float, default=0.0) # >0: init tok/pos std (GPT-standard 0.02)
args = ap.parse_args()
torch.manual_seed(args.seed)
dev = 'cuda' if torch.cuda.is_available() else 'cpu'
@@ -40,6 +41,9 @@ class Block(nn.Module):
tok = nn.Embedding(vocab, args.C).to(dev)
pos = nn.Embedding(args.T, args.C).to(dev)
+if args.tok_init > 0:
+ with torch.no_grad():
+ tok.weight.normal_(0, args.tok_init); pos.weight.normal_(0, args.tok_init)
blocks = nn.ModuleList([Block(args.C, args.H) for _ in range(args.L)]).to(dev)
mask = torch.triu(torch.full((args.T, args.T), float('-inf'), device=dev), 1)
params = list(tok.parameters()) + list(pos.parameters()) + list(blocks.parameters())