adaptive flag for unified trainer

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Nils
2026-07-14 12:32:36 +02:00
co-authored by Claude Fable 5
parent 32728e7b07
commit 95be35d20f
+5 -2
View File
@@ -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)