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:
@@ -15,7 +15,7 @@ from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from loop_common import BandLooper, MergeAdapter
|
||||
from loop_common import AdaptiveMergeAdapter, BandLooper, MergeAdapter
|
||||
from prep_mbpp import DIRECT_SUFFIX, extract_code, mbpp_prompt, run_tests
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
@@ -75,15 +75,16 @@ def main():
|
||||
ap.add_argument("--pause", type=int, default=0,
|
||||
help="append p pause tokens to each prompt")
|
||||
ap.add_argument("--alpha", type=float, default=0.3)
|
||||
ap.add_argument("--adaptive", action="store_true")
|
||||
args = ap.parse_args()
|
||||
ks = [int(x) for x in args.ks.split(",")]
|
||||
|
||||
model, tok = load_model(dtype=torch.bfloat16)
|
||||
tok.padding_side = "left"
|
||||
looper = BandLooper(model)
|
||||
adapter = MergeAdapter(
|
||||
d=model.config.get_text_config().hidden_size,
|
||||
alpha=args.alpha).cuda()
|
||||
cls = AdaptiveMergeAdapter if args.adaptive else MergeAdapter
|
||||
kw = ({"alpha0": args.alpha} if args.adaptive else {"alpha": args.alpha})
|
||||
adapter = cls(d=model.config.get_text_config().hidden_size, **kw).cuda()
|
||||
if args.adapter:
|
||||
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
|
||||
adapter.eval()
|
||||
|
||||
Reference in New Issue
Block a user