warm-start flag (stacking arm), adaptive flags for Blocksworld
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
+4
-2
@@ -11,7 +11,7 @@ import torch
|
||||
|
||||
from bw_common import verify_plan
|
||||
from bw_prep import DIRECT_SUFFIX, chat
|
||||
from loop_common import BandLooper, MergeAdapter
|
||||
from loop_common import AdaptiveMergeAdapter, BandLooper, MergeAdapter
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
from jlens.core import load_model # noqa: E402
|
||||
@@ -52,13 +52,15 @@ def main():
|
||||
ap.add_argument("--adapter", default=None)
|
||||
ap.add_argument("--tag", default="bw")
|
||||
ap.add_argument("--ks", default="0,2,4")
|
||||
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(
|
||||
cls = AdaptiveMergeAdapter if args.adaptive else MergeAdapter
|
||||
adapter = cls(
|
||||
d=model.config.get_text_config().hidden_size).cuda()
|
||||
if args.adapter:
|
||||
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
|
||||
|
||||
@@ -18,7 +18,7 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from bw_prep import DIRECT_SUFFIX, chat
|
||||
from loop_common import BandLooper, MergeAdapter
|
||||
from loop_common import AdaptiveMergeAdapter, BandLooper, MergeAdapter
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
from jlens.core import load_model # noqa: E402
|
||||
@@ -33,6 +33,7 @@ K_BUCKETS = [(1, ("easy",)), (2, ("easy", "hard")), (4, ("hard",))]
|
||||
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--seed", type=int, default=0)
|
||||
ap.add_argument("--adaptive", action="store_true")
|
||||
ARGS = ap.parse_args()
|
||||
MOVELIST_RE = re.compile(r"(?:^|\n)\s*1\s*[.)]", re.M)
|
||||
|
||||
@@ -81,7 +82,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)
|
||||
|
||||
@@ -127,7 +129,7 @@ def main():
|
||||
f"({(time.time()-t0)/(step+1):.1f}s/step)", flush=True)
|
||||
if step % 100 == 99 or step == STEPS - 1:
|
||||
torch.save(adapter.state_dict(),
|
||||
OUT / f"adapter_bw_e{step+1}.pt")
|
||||
OUT / f"adapter_bw{"_ad" if ARGS.adaptive else ""}_e{step+1}.pt")
|
||||
json.dump(log, open(OUT / "train_bw_log.json", "w"), indent=1)
|
||||
print("done", flush=True)
|
||||
|
||||
|
||||
@@ -51,6 +51,7 @@ ap.add_argument("--deepk", type=int, default=0,
|
||||
ap.add_argument("--bptt", type=int, default=0,
|
||||
help="truncated BPTT: grads only through last N iterations")
|
||||
ap.add_argument("--lr", type=float, default=1e-3)
|
||||
ap.add_argument("--warm", default=None, help="warm-start adapter checkpoint")
|
||||
ARGS, _ = ap.parse_known_args()
|
||||
SEED = ARGS.seed
|
||||
LR = ARGS.lr
|
||||
@@ -59,7 +60,8 @@ SUFFIX = ((f"_s{SEED}" if SEED else "")
|
||||
+ (f"_a{ARGS.alpha}" if ARGS.alpha != 0.3 else "")
|
||||
+ ("_ad" if ARGS.adaptive else "")
|
||||
+ (f"_dk{ARGS.deepk}" if ARGS.deepk else "")
|
||||
+ (f"_lr{ARGS.lr}" if ARGS.lr != 1e-3 else ""))
|
||||
+ (f"_lr{ARGS.lr}" if ARGS.lr != 1e-3 else "")
|
||||
+ ("_warm" if ARGS.warm else ""))
|
||||
PAUSE_ID = 6 # <unused0>
|
||||
|
||||
|
||||
@@ -122,6 +124,9 @@ def main():
|
||||
adapter = AdaptiveMergeAdapter(d=d, alpha0=ARGS.alpha).cuda()
|
||||
else:
|
||||
adapter = MergeAdapter(d=d, alpha=ARGS.alpha).cuda()
|
||||
if ARGS.warm:
|
||||
adapter.load_state_dict(torch.load(ARGS.warm, map_location="cuda"))
|
||||
print("warm-started from", ARGS.warm, flush=True)
|
||||
global K_BUCKETS
|
||||
if ARGS.deepk:
|
||||
f = ARGS.deepk / 4
|
||||
|
||||
Reference in New Issue
Block a user