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
|
||||||
import torch.nn.functional as F
|
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 DIRECT_SUFFIX as MBPP_SUFFIX
|
||||||
from prep_mbpp import mbpp_prompt
|
from prep_mbpp import mbpp_prompt
|
||||||
|
|
||||||
@@ -93,6 +94,7 @@ def main():
|
|||||||
help="feedforward control: adapter(e,e), no recurrence")
|
help="feedforward control: adapter(e,e), no recurrence")
|
||||||
ap.add_argument("--tasks", default="gsm,mbpp")
|
ap.add_argument("--tasks", default="gsm,mbpp")
|
||||||
ap.add_argument("--tag", default=None)
|
ap.add_argument("--tag", default=None)
|
||||||
|
ap.add_argument("--adaptive", action="store_true")
|
||||||
args = ap.parse_args()
|
args = ap.parse_args()
|
||||||
tasks = args.tasks.split(",")
|
tasks = args.tasks.split(",")
|
||||||
tag = args.tag or ("ff" if args.noloop else "uni")
|
tag = args.tag or ("ff" if args.noloop else "uni")
|
||||||
@@ -115,7 +117,8 @@ def main():
|
|||||||
for p in model.parameters():
|
for p in model.parameters():
|
||||||
p.requires_grad_(False)
|
p.requires_grad_(False)
|
||||||
looper = BandLooper(model)
|
looper = BandLooper(model)
|
||||||
adapter = MergeAdapter(
|
cls = AdaptiveMergeAdapter if args.adaptive else MergeAdapter
|
||||||
|
adapter = cls(
|
||||||
d=model.config.get_text_config().hidden_size).cuda()
|
d=model.config.get_text_config().hidden_size).cuda()
|
||||||
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
|
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user