adaptive-alpha merge (Lys-inspired, trained) + truncated-BPTT deep-k training
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -20,7 +20,7 @@ from pathlib import Path
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from loop_common import BandLooper, MergeAdapter
|
||||
from loop_common import AdaptiveMergeAdapter, BandLooper, MergeAdapter
|
||||
from prep_mbpp import DIRECT_SUFFIX, mbpp_prompt
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
@@ -34,6 +34,7 @@ LR = 1e-3
|
||||
WARMUP = 20
|
||||
MAX_TOK = 512
|
||||
K_BUCKETS = [(1, ("easy",)), (2, ("easy", "hard")), (4, ("hard",))]
|
||||
# --deepk 16 rescales to [(2, easy), (8, mixed), (16, hard)]
|
||||
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--seed", type=int, default=0)
|
||||
@@ -44,11 +45,19 @@ ap.add_argument("--feedforward", action="store_true",
|
||||
help="apply adapter once, no recurrence (pause-FF control)")
|
||||
ap.add_argument("--alpha", type=float, default=0.3,
|
||||
help="merge weight (2B-tuned default 0.3; try 0.1-0.15 at 12B)")
|
||||
ap.add_argument("--adaptive", action="store_true",
|
||||
help="state-dependent alpha (AdaptiveMergeAdapter)")
|
||||
ap.add_argument("--deepk", type=int, default=0,
|
||||
help="scale curriculum depths by deepk/4 (e.g. 16 -> 2/8/16)")
|
||||
ap.add_argument("--bptt", type=int, default=0,
|
||||
help="truncated BPTT: grads only through last N iterations")
|
||||
ARGS = ap.parse_args()
|
||||
SEED = ARGS.seed
|
||||
SUFFIX = ((f"_s{SEED}" if SEED else "")
|
||||
+ (f"_p{ARGS.pause}" if ARGS.pause else "")
|
||||
+ (f"_a{ARGS.alpha}" if ARGS.alpha != 0.3 else ""))
|
||||
+ (f"_a{ARGS.alpha}" if ARGS.alpha != 0.3 else "")
|
||||
+ ("_ad" if ARGS.adaptive else "")
|
||||
+ (f"_dk{ARGS.deepk}" if ARGS.deepk else ""))
|
||||
PAUSE_ID = 6 # <unused0>
|
||||
|
||||
|
||||
@@ -107,7 +116,15 @@ def main():
|
||||
p.requires_grad_(False)
|
||||
looper = BandLooper(model)
|
||||
d = model.config.get_text_config().hidden_size
|
||||
adapter = MergeAdapter(d=d, alpha=ARGS.alpha).cuda()
|
||||
if ARGS.adaptive:
|
||||
adapter = AdaptiveMergeAdapter(d=d, alpha0=ARGS.alpha).cuda()
|
||||
else:
|
||||
adapter = MergeAdapter(d=d, alpha=ARGS.alpha).cuda()
|
||||
global K_BUCKETS
|
||||
if ARGS.deepk:
|
||||
f = ARGS.deepk / 4
|
||||
K_BUCKETS = [(max(1, int(k * f)), lbls) for k, lbls in K_BUCKETS]
|
||||
print("K_BUCKETS ->", [(k, l) for k, l in K_BUCKETS], flush=True)
|
||||
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
|
||||
|
||||
train = [it for it in data if it["split"] == "train"
|
||||
@@ -135,7 +152,8 @@ def main():
|
||||
g["lr"] = lr_at(step)
|
||||
logits = looper.loop_logits(adapter, ids, k, attention_mask=msk,
|
||||
use_checkpoint=True, loop_mask=lmask,
|
||||
feedforward=ARGS.feedforward)
|
||||
feedforward=ARGS.feedforward,
|
||||
bptt=ARGS.bptt or None)
|
||||
loss = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(),
|
||||
lab[:, 1:].flatten(), ignore_index=-100)
|
||||
opt.zero_grad(set_to_none=True)
|
||||
|
||||
Reference in New Issue
Block a user