item 24 pre-registered: d=1 retry with loop-only band-LoRA r=16 (4.8M params, whole band, k=0 bit-exact); trainer/eval gain --bandlora; smoke-tested train+eval

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Nils
2026-07-16 21:44:01 +02:00
co-authored by Claude Fable 5
parent 7a141948fa
commit fe89820e64
4 changed files with 83 additions and 3 deletions
+33 -3
View File
@@ -53,10 +53,16 @@ ap.add_argument("--steps", type=int, default=STEPS)
ap.add_argument("--lr", type=float, default=LR)
ap.add_argument("--tag-suffix", default="",
help="appended to TAG (distinguish control variants)")
ap.add_argument("--bandlora", type=int, default=0, metavar="RANK",
help="item 24: loop-only LoRA (lora_band.LoopLoRA) on every "
"band layer, uniform scale 1.0 — active only during "
"band re-runs, k=0 stays bit-exact")
ap.add_argument("--lora-lr", type=float, default=1e-3)
ARGS = ap.parse_args()
STEPS, LR = ARGS.steps, ARGS.lr
TAG = ("carrycot_ff" if ARGS.feedforward else "carrycot") + (
f"_b{ARGS.drop_steps}" if ARGS.drop_steps else "") + (
f"_blr{ARGS.bandlora}" if ARGS.bandlora else "") + (
f"_ln{ARGS.lensnoise.replace(',', '_')}" if ARGS.lensnoise else "") + (
f"_s{ARGS.seed}" if ARGS.seed else "") + ARGS.tag_suffix
@@ -186,7 +192,25 @@ def main():
(adapter.base if ARGS.lensnoise else adapter).load_state_dict(
torch.load(ARGS.warm_start, map_location="cuda"))
print(f"warm-started from {ARGS.warm_start}", flush=True)
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
lora_params, lora_band_layers = [], []
if ARGS.bandlora:
from lora_band import inject_band_lora
from loop_common import BAND
lora_band_layers = list(range(BAND[0], BAND[1] + 1))
scales = {l: 1.0 for l in lora_band_layers}
lora_params = inject_band_lora(looper.tm, BAND[0], scales,
rank=ARGS.bandlora)
for p in lora_params:
p.data = p.data.cuda()
print(f"band-lora r={ARGS.bandlora}: "
f"{sum(p.numel() for p in lora_params)/1e6:.1f}M params, "
f"layers {lora_band_layers[0]}-{lora_band_layers[-1]}, "
f"lr={ARGS.lora_lr}", flush=True)
groups = [{"params": list(adapter.parameters()), "lr": LR, "base": LR}]
if lora_params:
groups.append({"params": lora_params, "lr": ARGS.lora_lr,
"base": ARGS.lora_lr})
opt = torch.optim.AdamW(groups, weight_decay=0.01)
keep = [it for it in data
if len(tok(it["question"])["input_ids"])
@@ -211,7 +235,7 @@ def main():
p = P_BY_LABEL[lbl]
ids, msk, lab, plens = build_batch(tok, batch, p)
for g in opt.param_groups:
g["lr"] = lr_at(step)
g["lr"] = g["base"] * lr_at(step) / LR
logits = carry_logits(looper, adapter, ids, msk, plens, K_PREFILL,
use_checkpoint=True,
feedforward=ARGS.feedforward)
@@ -219,7 +243,8 @@ def main():
lab[:, 1:].flatten(), ignore_index=-100)
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0)
torch.nn.utils.clip_grad_norm_(
list(adapter.parameters()) + lora_params, 1.0)
opt.step()
log.append({"step": step, "loss": loss.item()})
if step % 10 == 0:
@@ -235,6 +260,11 @@ def main():
adapter.noise_on = True
sd = (adapter.base if ARGS.lensnoise else adapter).state_dict()
torch.save(sd, OUT / f"adapter_{TAG}_e{step+1}.pt")
if lora_params:
torch.save({"rank": ARGS.bandlora, "band": lora_band_layers,
"tensors": [p.detach().cpu()
for p in lora_params]},
OUT / f"lora_{TAG}_e{step+1}.pt")
json.dump(log, open(OUT / f"train_{TAG}_log.json", "w"))
print("done", flush=True)