From 95be35d20ff557c4039d6a08aa03f7df8c0cb802 Mon Sep 17 00:00:00 2001 From: Nils Date: Tue, 14 Jul 2026 12:32:36 +0200 Subject: [PATCH] adaptive flag for unified trainer Co-Authored-By: Claude Fable 5 --- scripts/train_merge_unified.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/scripts/train_merge_unified.py b/scripts/train_merge_unified.py index df515fa..be1f39a 100644 --- a/scripts/train_merge_unified.py +++ b/scripts/train_merge_unified.py @@ -18,7 +18,8 @@ from pathlib import Path import torch import torch.nn.functional as F -from loop_common import BandLooper, MergeAdapter, chat_prompt, DIRECT_SUFFIX +from loop_common import (AdaptiveMergeAdapter, BandLooper, MergeAdapter, + chat_prompt, DIRECT_SUFFIX) from prep_mbpp import DIRECT_SUFFIX as MBPP_SUFFIX from prep_mbpp import mbpp_prompt @@ -93,6 +94,7 @@ def main(): help="feedforward control: adapter(e,e), no recurrence") ap.add_argument("--tasks", default="gsm,mbpp") ap.add_argument("--tag", default=None) + ap.add_argument("--adaptive", action="store_true") args = ap.parse_args() tasks = args.tasks.split(",") tag = args.tag or ("ff" if args.noloop else "uni") @@ -115,7 +117,8 @@ def main(): for p in model.parameters(): p.requires_grad_(False) looper = BandLooper(model) - adapter = MergeAdapter( + cls = AdaptiveMergeAdapter if args.adaptive else MergeAdapter + adapter = cls( d=model.config.get_text_config().hidden_size).cuda() opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)