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:
Nils
2026-07-14 10:29:57 +02:00
co-authored by Claude Fable 5
parent cd0ed80ea6
commit 53a8c9b609
3 changed files with 72 additions and 10 deletions
+5 -4
View File
@@ -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()