adaptive flag for unified trainer
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user