"""Go/no-go probe for the query-after-think garble recipe on Qwen3-8B (instruct).
Measures, BEFORE any training:
1. Competence: can the base instruct policy do state_track_q at all (credit at generous budget)?
-> need >~0.3 somewhere for GRPO signal; else SFT warm-start first.
2. Compression frontier: phase-A think generated ONCE at a generous cap, then TRUNCATED to each
budget B, force-closed with , and re-queried (phase B). credit-vs-B per T maps the
telegraphic-English capacity of the untrained policy -- the fixed budget the run should use is
one that stays feasible at low T and becomes the squeeze as T grows (the 1.5-3x band).
Truncate-and-requery is the right measurement: budgets are unobserved in the RL design, so the
policy cannot condition on B anyway.
3. No-think floor: forced-empty think -> the guess floor (load-bearing headroom).
4. Temperature density tail: P(correct AND non-prose-ish) by think temperature.
Run on the pod: python -m garble.probe_qfrontier --out /workspace/garble_runs/probe_qfrontier.json
"""
import argparse
import json
import os
import random
import statistics
try:
import dotenv; dotenv.load_dotenv()
except Exception:
pass
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
from transformers import AutoTokenizer # noqa: E402
from vllm import LLM, SamplingParams, TokensPrompt # noqa: E402
from garble.grpo_garble import garble_proxies, legibility_score # noqa: E402
from model_organisms.envs.base import SeqBuilder, initial_prefix_ids # noqa: E402
from model_organisms.envs.state_track import StateTrackQueryEnv # noqa: E402
TS = [6, 10, 14, 20, 28]
BUDGETS = [128, 192, 320, 512, 768, 1024, 2048]
TEMPS = [0.8, 1.0, 1.2, 1.4]
GEN_CAP = 2048
def phase_b_ids(tok, env, p, prefix, think_ids, empty=False):
sb = SeqBuilder(tok, prefix)
if empty:
sb.add_control("\n\n")
else:
sb.add_generated(list(think_ids))
sb.add_control("")
sb.close_assistant()
sb.add_user_turn(env.queries_text(p))
sb.add_control("\n\n\n\n")
return sb.ids
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--model", default="Qwen/Qwen3-8B")
ap.add_argument("--out", default="/workspace/garble_runs/probe_qfrontier.json")
ap.add_argument("--n-problems", type=int, default=12)
ap.add_argument("--k-samples", type=int, default=8)
ap.add_argument("--n-queries", type=int, default=3)
ap.add_argument("--gpu-mem", type=float, default=0.9)
args = ap.parse_args()
tok = AutoTokenizer.from_pretrained(args.model)
end_think_id = tok.encode("", add_special_tokens=False)[0]
im_end_id = tok.convert_tokens_to_ids("<|im_end|>")
llm = LLM(model=args.model, max_model_len=8192, gpu_memory_utilization=args.gpu_mem)
def gen(ids_list, max_tokens, temp, stop_ids):
sp = SamplingParams(max_tokens=max_tokens, temperature=temp, top_p=1.0, stop_token_ids=stop_ids)
outs = llm.generate([TokensPrompt(prompt_token_ids=i) for i in ids_list], sp)
res = []
for o in outs:
ids = list(o.outputs[0].token_ids)
while ids and ids[-1] in stop_ids:
ids = ids[:-1]
res.append(ids)
return res
results = {"grid": [], "nothink": [], "tempscan": []}
for T in TS:
env = StateTrackQueryEnv(r_min=4, r_max=4, t_min=T, t_max=T, val_max=30, k_max=9, mod=97,
n_queries=args.n_queries)
rng = random.Random(1000 + T)
probs = [env.sample_problem(rng) for _ in range(args.n_problems)]
prefixes = [initial_prefix_ids(tok, env.prompt(p)) for p in probs]
# phase A once per (problem, sample) at the generous cap
idx = [(pi, k) for pi in range(len(probs)) for k in range(args.k_samples)]
thinks = gen([prefixes[pi] for (pi, _k) in idx], GEN_CAP, 1.0, [end_think_id])
# no-think floor
nt_ids = [phase_b_ids(tok, env, p, pre, [], empty=True) for p, pre in zip(probs, prefixes)]
nt_ans = gen(nt_ids, 64, 0.6, [im_end_id])
nt_credit = [env.score_queries(p, tok.decode(a, skip_special_tokens=False))[0]
for p, a in zip(probs, nt_ans)]
results["nothink"].append({"T": T, "credit": sum(nt_credit) / len(nt_credit)})
# truncate-and-requery at each budget
for B in BUDGETS:
b_ids = [phase_b_ids(tok, env, probs[pi], prefixes[pi], th[:B])
for (pi, _k), th in zip(idx, thinks)]
answers = gen(b_ids, 64, 0.6, [im_end_id])
credits, alls, fit = [], [], []
correct_lens = []
for (pi, _k), th, a in zip(idx, thinks, answers):
c, _np = env.score_queries(probs[pi], tok.decode(a, skip_special_tokens=False))
credits.append(c)
alls.append(1.0 if c > 0.99 else 0.0)
fit.append(1.0 if len(th) <= B else 0.0)
if c > 0.99:
correct_lens.append(min(len(th), B))
cell = {"T": T, "B": B,
"credit": sum(credits) / len(credits),
"all_correct": sum(alls) / len(alls),
"fits": sum(fit) / len(fit),
"think_len_med": statistics.median(min(len(t), B) for t in thinks),
"correct_len_min": min(correct_lens) if correct_lens else None,
"correct_len_med": statistics.median(correct_lens) if correct_lens else None}
results["grid"].append(cell)
print(f"[grid] T={T:>2} B={B:>4} credit={cell['credit']:.3f} all={cell['all_correct']:.3f} "
f"fits={cell['fits']:.2f} lenmed={cell['think_len_med']:.0f}", flush=True)
# temperature density tail at T=10, generous budget
env = StateTrackQueryEnv(r_min=4, r_max=4, t_min=10, t_max=10, val_max=30, k_max=9, mod=97,
n_queries=args.n_queries)
rng = random.Random(1010)
probs = [env.sample_problem(rng) for _ in range(args.n_problems)]
prefixes = [initial_prefix_ids(tok, env.prompt(p)) for p in probs]
idx = [(pi, k) for pi in range(len(probs)) for k in range(args.k_samples)]
for temp in TEMPS:
thinks = gen([prefixes[pi] for (pi, _k) in idx], GEN_CAP, temp, [end_think_id])
b_ids = [phase_b_ids(tok, env, probs[pi], prefixes[pi], th) for (pi, _k), th in zip(idx, thinks)]
answers = gen(b_ids, 64, 0.6, [im_end_id])
rows = []
for (pi, _k), th, a in zip(idx, thinks, answers):
c, _np = env.score_queries(probs[pi], tok.decode(a, skip_special_tokens=False))
text = tok.decode(th, skip_special_tokens=False)
px = garble_proxies(text)
rows.append({"credit": c, "len": len(th), "leg": legibility_score(text),
"non_ascii": px["non_ascii"], "alpha": px["alpha_word_ratio"]})
ok = [r for r in rows if r["credit"] > 0.99]
cell = {"temp": temp, "credit": sum(r["credit"] for r in rows) / len(rows),
"all_correct": len(ok) / len(rows),
"leg_mean": sum(r["leg"] for r in rows) / len(rows),
"leg_correct": (sum(r["leg"] for r in ok) / len(ok)) if ok else None,
"nonprose_and_correct": sum(1 for r in ok if r["alpha"] < 0.35) / len(rows),
"len_med": statistics.median(r["len"] for r in rows)}
results["tempscan"].append(cell)
print(f"[temp] T=10 temp={temp} credit={cell['credit']:.3f} all={cell['all_correct']:.3f} "
f"leg={cell['leg_mean']:.3f} nonprose&ok={cell['nonprose_and_correct']:.3f} "
f"lenmed={cell['len_med']:.0f}", flush=True)
os.makedirs(os.path.dirname(args.out), exist_ok=True)
with open(args.out, "w") as f:
json.dump(results, f, indent=2)
print(f"[done] wrote {args.out}", flush=True)
if __name__ == "__main__":
main()