128 lines
6.7 KiB
Python
128 lines
6.7 KiB
Python
"""Entrypoint: emergent-garble GRPO on Qwen3-8B (instruct/thinking) with the QUERY-AFTER-THINK
|
|
register env -- the composed recipe:
|
|
|
|
* state_track_q: random (op, register) queries revealed only AFTER </think>, answered with a
|
|
forced-empty think -> the whole trajectory must be decodably present in the CoT (load-bearing
|
|
by construction, no lucky-short-guess mode).
|
|
* Lane-2 squeeze: T (ops) GROWS across tiers while max_think stays FIXED -> past the point where
|
|
verbose prose fits, density is required by the ACCURACY channel itself (nothing to "eat").
|
|
Dwell-gated tier advance; adaptive-lambda length penalty (on think tokens, correct-only) does
|
|
the fine-grained ordering among fitting rollouts.
|
|
* Rank tournament (rank_bonus): zero-mean inverse-length-rank bonus among fully-correct rollouts
|
|
keeps the compression gradient alive when lengths cluster (the rs5 gradient-death fix).
|
|
* Full-FT (LoRA caps the mode shift), asymmetric temps (hot think, cold answers), no KL anchor
|
|
on the think region (the instruct prior IS the anti-drift force we're releasing).
|
|
|
|
Pure RL + length penalty; no SFT, no legibility/monitor terms, no vocab constraints.
|
|
|
|
python -m garble.train_garble_q8b --batch-id gq1 --run-name gq1_t512 --max-think 512
|
|
"""
|
|
import os
|
|
|
|
try:
|
|
import dotenv; dotenv.load_dotenv()
|
|
except Exception:
|
|
pass
|
|
os.environ.setdefault("HF_HOME", f"/workspace-vast/{os.environ.get('USER', 'jbauer')}/hf_cache")
|
|
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
|
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
|
|
|
|
import argparse # noqa: E402
|
|
import json # noqa: E402
|
|
|
|
from garble.grpo_garble import train # noqa: E402 (vanilla HF full-FT; NO unsloth on purpose)
|
|
|
|
|
|
def t_ladder(r, ts, n_queries, mod=97, val_max=30):
|
|
return [{"r_min": r, "r_max": r, "t_min": t, "t_max": t + 2, "val_max": val_max, "k_max": 9,
|
|
"mod": mod, "n_queries": n_queries} for t in ts]
|
|
|
|
|
|
def entity_ladder(e, ts, n_queries):
|
|
return [{"e_min": e, "e_max": e, "t_min": t, "t_max": t + 2, "n_queries": n_queries} for t in ts]
|
|
|
|
|
|
DEFAULTS = dict(
|
|
model_name="Qwen/Qwen3-8B", # instruct/thinking -- Base degenerates at 8B (word-salad, not garble)
|
|
full_ft=True, use_vllm=False, base_model=False, grad_ckpt=True,
|
|
env="state_track_q",
|
|
# Lane 2: T grows, budget fixed. Tiers t_min..t_min+2 for within-tier variety.
|
|
curriculum=t_ladder(4, [6, 9, 12, 16, 20, 24, 28, 34, 40, 48], n_queries=3),
|
|
auto_curriculum=True, start_tier=0, tier_advance_thresh=0.75, tier_min_dwell=8,
|
|
lr=2e-6, warmup_ratio=0.02, max_grad_norm=1.0, max_steps=600,
|
|
batch_problems=8, group_size=12, # 96 rollouts/step, 8 GRPO groups
|
|
pg_microbatch=2, gen_chunk=48,
|
|
temperature=1.0, top_p=1.0, think_temperature=1.1, out_temperature=0.6,
|
|
max_think=1024, max_out=64, max_seq_length=8192,
|
|
# reward: query-credit - lambda*(think excess); correct = ALL queries right (partial credit
|
|
# still flows through task_reward; the penalty/tournament keys on fully-correct only)
|
|
correct_thresh=0.99, rank_bonus=0.25,
|
|
adaptive_lambda=True, lambda_target=0.60, lambda_lr=0.02, lambda_ema_alpha=0.9,
|
|
lambda_init=0.0, lambda_max=1.5, lambda_penalty=0.3,
|
|
length_penalty_mode="linear", length_target=None, length_floor=None, length_anneal_steps=0,
|
|
penalty_signal="length", penalty_on_correct_only=True, w_len=0.5, w_leg=1.0,
|
|
normalize_advantages=True,
|
|
eval_every=10, eval_problems=16,
|
|
seed=42, n_checkpoints=12, # 8B full-FT ckpts are 16GB each -- keep the geometric set lean
|
|
keep_ckpts=3, latest_optimizer=True, train_answer_tokens=True,
|
|
save_dir="/root/garble_runs", # pod overlay (500G, fast); /workspace quota is ~50G on this pod
|
|
wandb_project="cot-oracle", wandb_entity="MATS10-CS-JB",
|
|
wandb_group="", run_name="", console_every=1, table_every=10, save_every=25, resume_from="",
|
|
lora_r=0, lora_alpha=0, lora_dropout=0.0, # unused in full-FT
|
|
)
|
|
|
|
|
|
def main():
|
|
ap = argparse.ArgumentParser()
|
|
ap.add_argument("--batch-id", default=None)
|
|
ap.add_argument("--run-name", default=None)
|
|
ap.add_argument("--env", default=None, choices=["state_track_q", "entity_track_q"])
|
|
ap.add_argument("--model-name", default=None)
|
|
ap.add_argument("--max-steps", type=int, default=None)
|
|
ap.add_argument("--max-think", type=int, default=None)
|
|
ap.add_argument("--lr", type=float, default=None)
|
|
ap.add_argument("--lambda-target", type=float, default=None)
|
|
ap.add_argument("--rank-bonus", type=float, default=None)
|
|
ap.add_argument("--think-temperature", type=float, default=None)
|
|
ap.add_argument("--start-tier", type=int, default=None)
|
|
ap.add_argument("--tier-min-dwell", type=int, default=None)
|
|
ap.add_argument("--n-queries", type=int, default=None)
|
|
ap.add_argument("--keep-ckpts", type=int, default=None)
|
|
ap.add_argument("--latest-optimizer", type=int, default=None)
|
|
ap.add_argument("--train-answer-tokens", type=int, default=None)
|
|
ap.add_argument("--length-penalty-mode", default=None, choices=["linear", "target_excess", "anneal_excess"])
|
|
ap.add_argument("--length-target", type=int, default=None)
|
|
ap.add_argument("--length-floor", type=int, default=None)
|
|
ap.add_argument("--length-anneal-steps", type=int, default=None)
|
|
ap.add_argument("--seed", type=int, default=None)
|
|
ap.add_argument("--save-dir", default=None)
|
|
ap.add_argument("--resume-from", default=None)
|
|
ap.add_argument("--wandb-mode", default=None, choices=["online", "offline", "disabled"])
|
|
args = ap.parse_args()
|
|
if args.wandb_mode:
|
|
os.environ["WANDB_MODE"] = args.wandb_mode
|
|
cfg = dict(DEFAULTS)
|
|
cli = {k: getattr(args, k) for k in ("model_name", "max_steps", "max_think", "lr", "lambda_target",
|
|
"rank_bonus", "think_temperature", "start_tier", "tier_min_dwell", "length_penalty_mode",
|
|
"length_target", "length_floor", "length_anneal_steps", "seed", "save_dir", "run_name",
|
|
"resume_from", "keep_ckpts", "latest_optimizer", "train_answer_tokens")}
|
|
cfg.update({k: v for k, v in cli.items() if v is not None})
|
|
if args.env:
|
|
cfg["env"] = args.env
|
|
if args.env == "entity_track_q":
|
|
cfg["curriculum"] = entity_ladder(5, [6, 9, 12, 16, 20, 24, 28, 34, 40, 48], n_queries=3)
|
|
if args.n_queries is not None:
|
|
for tier in cfg["curriculum"]:
|
|
tier["n_queries"] = args.n_queries
|
|
if args.batch_id:
|
|
cfg["wandb_group"] = args.batch_id
|
|
if not cfg["run_name"]:
|
|
cfg["run_name"] = f"gq_{args.batch_id or 'x'}_{cfg['seed']}"
|
|
print("[config]", json.dumps({k: v for k, v in cfg.items() if k != "curriculum"}, indent=2), flush=True)
|
|
print("[curriculum]", json.dumps(cfg["curriculum"]), flush=True)
|
|
train(cfg)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|