From 32d1fdbfd0becef18351762d5d1d2924640a52be Mon Sep 17 00:00:00 2001 From: Yuren Hao Date: Thu, 9 Jul 2026 09:52:33 -0500 Subject: 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 Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn --- ep_run/casc_bp_train.py | 4 ++++ 1 file changed, 4 insertions(+) (limited to 'ep_run/casc_bp_train.py') 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()) -- cgit v1.2.3