RecurrentAdapter arm: Huginn-regime retrofit (learned A/B, noise h0, randomized depth) + pre-registration item 11

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Nils
2026-07-15 00:49:44 +02:00
co-authored by Claude Fable 5
parent 0fd93cb328
commit d2da8044c3
4 changed files with 94 additions and 9 deletions
+6 -2
View File
@@ -15,7 +15,8 @@ from pathlib import Path
import torch
from loop_common import AdaptiveMergeAdapter, BandLooper, MergeAdapter
from loop_common import (AdaptiveMergeAdapter, BandLooper, MergeAdapter,
RecurrentAdapter)
from prep_mbpp import DIRECT_SUFFIX, extract_code, mbpp_prompt, run_tests
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
@@ -76,13 +77,16 @@ def main():
help="append p pause tokens to each prompt")
ap.add_argument("--alpha", type=float, default=0.3)
ap.add_argument("--adaptive", action="store_true")
ap.add_argument("--rec", action="store_true",
help="RecurrentAdapter (noise h0, learned A/B)")
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)
cls = AdaptiveMergeAdapter if args.adaptive else MergeAdapter
cls = (RecurrentAdapter if args.rec
else 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: