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:
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user