"""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 , 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()