133 lines
5.0 KiB
Markdown
133 lines
5.0 KiB
Markdown
---
|
|
license: apache-2.0
|
|
base_model: Qwen/Qwen2.5-1.5B-Instruct
|
|
library_name: transformers
|
|
pipeline_tag: text-generation
|
|
language:
|
|
- en
|
|
tags:
|
|
- grpo
|
|
- rlvr
|
|
- trl
|
|
- reasoning
|
|
- 24-game
|
|
- countdown
|
|
- verifier-guided-sampling
|
|
---
|
|
|
|
# Qwen2.5-1.5B-Instruct · GRPO/RLVR for the 24 Game
|
|
|
|
GRPO-trained `Qwen/Qwen2.5-1.5B-Instruct` that solves the **24 game** (and, by
|
|
transfer, **Countdown**) in R1 style. Rewards come from a **deterministic program
|
|
verifier** (RLVR) — no neural reward model. This is the `shaped`-reward GRPO
|
|
checkpoint (300 steps, from an SFT warmup).
|
|
|
|
Output format:
|
|
|
|
```text
|
|
<think>reasoning</think><answer>3*8</answer>
|
|
```
|
|
|
|
The `<answer>` may only contain the input numbers, `+ - * /`, and parentheses;
|
|
each input number is used exactly once; the value must equal the target (24 by
|
|
default, arbitrary for Countdown).
|
|
|
|
## TL;DR — how to use it
|
|
|
|
> **This model is a verifier-guided *sampler*, not a reliable one-shot solver.**
|
|
> Greedy/pass@1 ≈ 0. Its value is that **sampling many candidates and letting the
|
|
> verifier pick a correct one** works very well. Use temperature ~1.1, sample N,
|
|
> then run the program verifier to select the first valid expression.
|
|
|
|
## Results (100 hardest problems, human-lowest-solved-rate split)
|
|
|
|
| Metric | This model | Note |
|
|
| --- | ---: | --- |
|
|
| greedy / pass@1 | **0/100** | GRPO does not improve one-shot accuracy |
|
|
| sample@16 (per independent run) | **~36/100** (max 43) | temp 1.1, top-p 0.95 |
|
|
| best-of-128 (single model, verifier-selected) | **74/100** | 128 candidates/problem |
|
|
| union best-of-runs (this + correctness-heavy, 256/problem) | **79/100** | two GRPO checkpoints merged |
|
|
| **adaptive best-of-runs** | **100/100** | escalate sampling only on unsolved; 60,416 total candidates, hardest problem needed 5,376 |
|
|
|
|
Diagnostics on first attempts: `r1_format = 1.00`, `numbers_ok = 0.99` — the model
|
|
reliably learned the format and exact number-usage constraint, but not exact
|
|
arithmetic search. That is why test-time sampling + verifier selection is the
|
|
right way to use it.
|
|
|
|
### Countdown (bonus, arbitrary target)
|
|
|
|
| Setting | greedy | sample@16 |
|
|
| --- | ---: | ---: |
|
|
| this 24-game checkpoint, zero-shot | 1/100 | 11/100 |
|
|
| after 50-step target-aware Countdown GRPO | 1/100 | 18/100 |
|
|
|
|
## Usage
|
|
|
|
```python
|
|
import re
|
|
from fractions import Fraction
|
|
from transformers import AutoModelForCausalLM, AutoTokenizer
|
|
import torch
|
|
|
|
name = "lxazjk/qwen2.5-1.5b-24game-grpo"
|
|
tok = AutoTokenizer.from_pretrained(name)
|
|
model = AutoModelForCausalLM.from_pretrained(name, torch_dtype="auto", device_map="auto")
|
|
|
|
numbers = [2, 4, 7, 7]
|
|
prompt = (
|
|
"You solve arithmetic target games. Use each given number exactly once. "
|
|
"Only use +, -, *, / and parentheses. Think briefly, then put the final "
|
|
"expression inside <answer>...</answer>. The answer must be only an "
|
|
"expression, with no equals sign, result value, example, or explanation.\n\n"
|
|
f"Make 24.\nNumbers: {', '.join(map(str, numbers))}\n"
|
|
"Respond in this exact format:\n<think>reasoning steps</think><answer>expression only</answer>"
|
|
)
|
|
|
|
# sample many candidates, then verify
|
|
ids = tok(prompt, return_tensors="pt").to(model.device)
|
|
out = model.generate(**ids, do_sample=True, temperature=1.1, top_p=0.95,
|
|
num_return_sequences=16, max_new_tokens=128)
|
|
cands = tok.batch_decode(out[:, ids.input_ids.shape[1]:], skip_special_tokens=True)
|
|
|
|
def verify(expr, nums, target=24):
|
|
toks = re.findall(r"\d+", expr)
|
|
if sorted(map(int, toks)) != sorted(nums):
|
|
return False
|
|
if re.fullmatch(r"[0-9+\-*/() ]+", expr) is None:
|
|
return False
|
|
try:
|
|
# evaluate with exact rational arithmetic
|
|
val = eval(expr, {"__builtins__": {}})
|
|
return Fraction(val) == target
|
|
except Exception:
|
|
return False
|
|
|
|
for c in cands:
|
|
m = re.search(r"<answer>(.*?)</answer>", c, re.S)
|
|
if m and verify(m.group(1).strip(), numbers):
|
|
print("solved:", m.group(1).strip()); break
|
|
```
|
|
|
|
(In the original project the verifier uses a safe AST whitelist + `Fraction`; the
|
|
snippet above is a compact stand-in. Never `eval` untrusted strings in production.)
|
|
|
|
## Training
|
|
|
|
- Base: `Qwen/Qwen2.5-1.5B-Instruct`
|
|
- SFT warmup (format + number usage), then **GRPO** with a rule-based reward set
|
|
(`format / structure / expression / number-usage / closeness / correctness`).
|
|
- Verifier: extract last `<answer>`, restrict characters, check number multiset,
|
|
parse via a safe AST whitelist, evaluate with `Fraction`, compare to target.
|
|
- Data: `nlile/24-game` (train), `test-time-compute/game-of-24` (hard eval),
|
|
`Jiayi-Pan/Countdown-Tasks-3to4` (bonus).
|
|
|
|
See `results/` in this repo for the full per-setting evaluation summaries
|
|
(`REPRODUCTION_SUMMARY.md`, greedy / sample / union / adaptive / Countdown JSONs).
|
|
|
|
## Limitations
|
|
|
|
- Not a one-shot solver — high solve rates require substantial test-time sampling.
|
|
- The hard-eval number combinations overlap the training set, so this is a hard
|
|
eval, not a strict OOD generalization test.
|
|
- `100/100` is a test-time sampling + verifier-selection result, **not** pass@1.
|