Verifiable outcome × Formal Reasoning Engines × DeepSeek-V3.1. Loss importance_sampling. The compiled file is in the page source.
RL-tune DeepSeek-V3.1 vs a Lean kernel on Tinker
formal-reasoning-deepseek-v3-1-environment.pysample · forward_backward · optim_step · save_state
"""Reinforcement.tech compiled Tinker loop
H1: Reinforcement: Build Your Own Reward Model
INPUT
Signal : Verifiable outcome (Environment RL)
In : A verifier — Lean, pytest, compiler, retrieval set.
Task : Formal Reasoning Engines — Checkable traces
Model : deepseek-ai/DeepSeek-V3.1 (MOE, large MoE)
OUTPUT
Loop : Runnable environment loop.
Loss : importance_sampling
LoRA rank : 32
Steps : 50
Tinker primitives used in this file:
sample generate on-policy rollouts
forward_backward accumulate LoRA gradients
optim_step Adam update on the adapter
save_state checkpoint weights + optimizer
Requires:
uv pip install tinker pytest
export TINKER_API_KEY=...
Docs: https://tinker-docs.thinkingmachines.ai/tinker/quickstart/
Dataset JSONL (one record per line). Pass --data PATH.
environment : {"prompt": str, "metadata": {...}}
dpo : {"prompt": str, "chosen": str, "rejected": str}
sdft : {"prompt": str, "completion": str}
metadata is passed through to the verifier (expected answer, test file, Lean goal).
--eval-data uses the same schema on a held-out set (never optim_step).
Without --data this file runs one built-in example (a smoke test, not training).
CLI (overrides the module constants; --seed is only an args field):
--data --eval-data --eval-every --log --steps --rank --lr --out --seed
--resume --group-size --prompts-per-step --max-tokens --temperature
"""
from __future__ import annotations
import argparse
import asyncio
import json
import os
import random
from pathlib import Path
try:
import tinker
from tinker import types
except ImportError: # --help and load_dataset work without Tinker installed
tinker = None
types = None
BASE_MODEL = "deepseek-ai/DeepSeek-V3.1"
LORA_RANK = 32
LEARNING_RATE = 5e-4
STEPS = 50
MAX_TOKENS = 256
TEMPERATURE = 0.8
GROUP_SIZE = 8
PROMPTS_PER_STEP = 1
EVAL_EVERY = 10
SAVE_EVERY = 10
RUN_NAME = Path(__file__).stem
DATASET_KEYS = ("prompt",)
LARGE_MODEL_NOTE = "warning: LEARNING_RATE=5e-4 is the MoE default. DeepSeek-V3.1 is large MoE, so you may want it smaller. Pass --lr."
def parse_args():
parser = argparse.ArgumentParser(description="Reinforcement.tech compiled Tinker loop")
parser.add_argument("--data", help="JSONL dataset path")
parser.add_argument("--eval-data", help="Held-out JSONL. Same schema as --data. Never used for optim_step.")
parser.add_argument("--eval-every", type=int, default=EVAL_EVERY)
parser.add_argument("--log", help="Append per-step JSONL: {step, mean_reward, n_datums, loss}")
parser.add_argument("--steps", type=int, default=STEPS)
parser.add_argument("--rank", type=int, default=LORA_RANK)
parser.add_argument("--lr", type=float, default=float(LEARNING_RATE))
parser.add_argument("--out", default=RUN_NAME, help="Run name / checkpoint prefix")
parser.add_argument(
"--seed",
type=int,
default=0,
help="Shuffle seed for --data / --eval-data. Not copied in apply_args; rows_for_run reads args.seed.",
)
parser.add_argument("--resume", help="tinker:// path printed by save_state (weights + optimizer)")
parser.add_argument("--group-size", type=int, default=GROUP_SIZE)
parser.add_argument("--prompts-per-step", type=int, default=PROMPTS_PER_STEP)
parser.add_argument("--max-tokens", type=int, default=MAX_TOKENS)
parser.add_argument("--temperature", type=float, default=TEMPERATURE)
return parser.parse_args()
def apply_args(args) -> None:
global LORA_RANK, STEPS, LEARNING_RATE, RUN_NAME
global GROUP_SIZE, PROMPTS_PER_STEP, MAX_TOKENS, TEMPERATURE, EVAL_EVERY
# --seed is intentionally not a module constant. rows_for_run(args) and the
# eval shuffle read args.seed so two shuffles stay independent of apply_args.
LORA_RANK = args.rank
STEPS = args.steps
LEARNING_RATE = args.lr
RUN_NAME = args.out
GROUP_SIZE = args.group_size
PROMPTS_PER_STEP = args.prompts_per_step
MAX_TOKENS = args.max_tokens
TEMPERATURE = args.temperature
EVAL_EVERY = args.eval_every
def require_key() -> None:
if not os.environ.get("TINKER_API_KEY"):
raise SystemExit("Set TINKER_API_KEY before running this loop.")
def load_dataset(path: str) -> list[dict]:
"""Read JSONL. One record per line. Schema: see the module docstring."""
rows: list[dict] = []
with open(path, encoding="utf-8") as handle:
for line_no, line in enumerate(handle, 1):
line = line.strip()
if not line:
continue
row = json.loads(line)
missing = [key for key in DATASET_KEYS if key not in row]
if missing:
raise SystemExit(f"{path}:{line_no} missing {missing}")
rows.append(row)
if not rows:
raise SystemExit(f"{path} had no JSONL rows.")
return rows
def rows_for_run(args) -> list[dict]:
if args.data:
rows = load_dataset(args.data)
else:
print("No --data given; running on 1 built-in example. This is a smoke test, not training.")
rows = list(EXAMPLE_ROWS)
rng = random.Random(args.seed)
rng.shuffle(rows)
needed = max(args.steps, 1) * max(args.prompts_per_step, 1)
if needed > len(rows):
print(f"warning: --steps/--prompts-per-step need {needed} rows; dataset has {len(rows)}; cycling.")
return rows
def eval_rows_for_run(args) -> list[dict]:
if not args.eval_data:
return []
rows = load_dataset(args.eval_data)
rng = random.Random(args.seed)
rng.shuffle(rows)
return rows
def print_config(args, rows: list[dict], eval_rows: list[dict]) -> None:
print("=== run config ===")
print(f"model={BASE_MODEL}")
print(f"rank={LORA_RANK} steps={STEPS} lr={LEARNING_RATE} seed={args.seed}")
print(f"group_size={GROUP_SIZE} prompts_per_step={PROMPTS_PER_STEP}")
print(f"max_tokens={MAX_TOKENS} temperature={TEMPERATURE}")
print(f"data={args.data or '(EXAMPLE_ROWS)'} n_train={len(rows)}")
print(f"eval_data={args.eval_data or '(none)'} n_eval={len(eval_rows)} eval_every={EVAL_EVERY}")
print(f"log={args.log or '(none)'} resume={args.resume or '(none)'} out={RUN_NAME}")
print("==================")
if LARGE_MODEL_NOTE:
print(LARGE_MODEL_NOTE)
def append_metrics(path: str | None, record: dict) -> None:
if not path:
return
out = Path(path)
out.parent.mkdir(parents=True, exist_ok=True)
with out.open("a", encoding="utf-8") as handle:
handle.write(json.dumps(record) + "\n")
def metric_loss(result) -> float | None:
metrics = getattr(result, "metrics", None) or {}
if hasattr(metrics, "get"):
for key in ("loss", "importance_sampling"):
if metrics.get(key) is not None:
return metrics[key]
loss = getattr(result, "loss", None)
return float(loss) if loss is not None else None
async def connect(args):
if tinker is None or types is None:
raise SystemExit("uv pip install tinker — then rerun.")
service = tinker.ServiceClient()
if args.resume:
training = await service.create_training_client_from_state_with_optimizer_async(args.resume)
print(f"resumed weights+optimizer from {args.resume}")
else:
training = await service.create_lora_training_client_async(
base_model=BASE_MODEL,
rank=LORA_RANK,
user_metadata={"product": "reinforcement.tech", "run": RUN_NAME},
)
tokenizer = training.get_tokenizer()
return service, training, tokenizer
async def sampling_client(training):
"""Ephemeral on-policy sampler. Do not pass name= — it is deprecated and ignored."""
return await training.save_weights_and_get_sampling_client_async()
async def checkpoint(training, step: int) -> None:
if STEPS == 0:
return
if step % SAVE_EVERY != 0 and step != STEPS - 1:
return
saved = await training.save_state_async(name=f"{RUN_NAME}-step-{step}")
result = await saved.result_async()
path = getattr(result, "path", None)
print(f"checkpoint name={RUN_NAME}-step-{step} path={path}")
print("resume later with --resume PATH")
PROMPT = "Prove the statement in Lean 4. Return only a complete proof."
# Built-in smoke-test row. Used only when --data is absent.
EXAMPLE_ROWS = [{"prompt": PROMPT, "metadata": {}}]
# Pattern: Checkable traces. Renderer family hint: deepseek.
def verify_environment(completion: str, metadata: dict | None = None) -> float:
"""Lean-shaped stub. A mediocre completion must fail.
Require a theorem/lemma declaration, a `:= by` / `begin` body, and no
sorry/admit. Swap this for verify_with_pytest or `lean --make`.
"""
import re
has_decl = re.search(r"\b(theorem|lemma)\b", completion, flags=re.I) is not None
has_body = re.search(r"(:=\s*by|\bbegin\b)", completion, flags=re.I) is not None
has_hole = re.search(r"\b(sorry|admit)\b", completion, flags=re.I) is not None
return 1.0 if has_decl and has_body and not has_hole else 0.0
def verify_with_pytest(completion: str, timeout: float = 8.0) -> float:
"""Real checker: write the completion and run pytest -q."""
import subprocess
import tempfile
from pathlib import Path
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "test_completion.py"
path.write_text(completion, encoding="utf-8")
try:
proc = subprocess.run(
["pytest", "-q", str(path)],
timeout=timeout,
capture_output=True,
check=False,
)
except (FileNotFoundError, subprocess.TimeoutExpired):
return 0.0
return 1.0 if proc.returncode == 0 else 0.0
def pack_rl_datum(
tokenizer,
prompt: str,
completion_tokens: list[int],
sampling_logprobs: list[float],
advantage: float,
) -> types.Datum:
prompt_tokens = tokenizer.encode(prompt)
full = prompt_tokens + completion_tokens
n_prefix = max(len(prompt_tokens) - 1, 0)
return types.Datum(
model_input=types.ModelInput.from_ints(tokens=full[:-1]),
loss_fn_inputs=dict(
target_tokens=full[1:],
logprobs=[0.0] * n_prefix + list(sampling_logprobs),
advantages=[0.0] * n_prefix + [advantage] * len(completion_tokens),
),
)
def rollout_logprobs(sequence) -> list[float]:
logprobs = sequence.logprobs if getattr(sequence, "logprobs", None) else None
if logprobs is None:
raise RuntimeError(
"Sampler returned no logprobs. The importance-sampling loss requires "
"per-token logprobs from the sampling client. Check your SamplingParams."
)
return list(logprobs)
NORMALIZE_ADVANTAGE = False # True → divide (reward − mean) by group std
# Advantage is reward − mean, computed inside each prompt's group — not across
# the whole step. We do not divide by std (NORMALIZE_ADVANTAGE) and we do not
# add a KL term against a frozen reference. This is on-policy IS versus the
# sampler that produced the tokens, not textbook PPO.
def advantages(rewards: list[float]) -> list[float]:
baseline = sum(rewards) / max(len(rewards), 1)
if len(rewards) > 1:
var = sum((reward - baseline) ** 2 for reward in rewards) / len(rewards)
std = var ** 0.5
else:
std = 0.0
out = [reward - baseline for reward in rewards]
if NORMALIZE_ADVANTAGE and std > 1e-8:
out = [value / std for value in out]
return out
async def rollout_prompt(sampling, tokenizer, row, params):
prompt_text = str(row["prompt"])
prompt = types.ModelInput.from_ints(tokenizer.encode(prompt_text))
rollout = await sampling.sample_async(
prompt=prompt,
num_samples=GROUP_SIZE,
sampling_params=params,
)
rewards = [
verify_environment(tokenizer.decode(sequence.tokens), row.get("metadata") or {})
for sequence in rollout.sequences
]
baseline = sum(rewards) / max(len(rewards), 1)
if all(reward == rewards[0] for reward in rewards):
return prompt_text, rewards, [], baseline, True
adv = advantages(rewards)
datums = [
pack_rl_datum(
tokenizer,
prompt_text,
sequence.tokens,
rollout_logprobs(sequence),
advantage,
)
for sequence, advantage in zip(rollout.sequences, adv)
]
return prompt_text, rewards, datums, baseline, False
async def evaluate(eval_rows, training, tokenizer, params) -> float:
sampling = await sampling_client(training)
rewards: list[float] = []
for row in eval_rows:
prompt = types.ModelInput.from_ints(tokenizer.encode(str(row["prompt"])))
rollout = await sampling.sample_async(prompt=prompt, num_samples=1, sampling_params=params)
rewards.append(
verify_environment(tokenizer.decode(rollout.sequences[0].tokens), row.get("metadata") or {})
)
return sum(rewards) / max(len(rewards), 1)
async def train(args) -> None:
require_key()
rows = rows_for_run(args)
eval_rows = eval_rows_for_run(args)
print_config(args, rows, eval_rows)
_service, training, tokenizer = await connect(args)
params = types.SamplingParams(max_tokens=MAX_TOKENS, temperature=TEMPERATURE)
consecutive_skips = 0
preview_prompt_text = str(rows[0]["prompt"]) if rows else PROMPT
for step in range(STEPS):
sampling = await sampling_client(training)
step_rewards: list[float] = []
step_datums: list[types.Datum] = []
skipped_groups = 0
for offset in range(PROMPTS_PER_STEP):
row = rows[(step * PROMPTS_PER_STEP + offset) % len(rows)]
_text, rewards, datums, baseline, skipped = await rollout_prompt(
sampling, tokenizer, row, params
)
step_rewards.extend(rewards)
if skipped:
skipped_groups += 1
if consecutive_skips == 0 and skipped_groups == 1:
print(
"WARNING: every sample scored "
f"{rewards[0]:.3f}. verify_environment is not separating "
"completions, so this group skips optim_step. Replace "
"the placeholder so some samples score 0 and others score 1."
)
print(f"step {step:03d} skip degenerate group R={baseline:.3f}")
else:
step_datums.extend(datums)
if skipped_groups == PROMPTS_PER_STEP:
consecutive_skips += 1
mean_reward = sum(step_rewards) / max(len(step_rewards), 1)
append_metrics(args.log, {"step": step, "mean_reward": mean_reward, "n_datums": 0, "loss": None})
if consecutive_skips >= 5:
raise SystemExit(
"Stopped after 5 consecutive degenerate groups. "
"verify_environment is not separating samples. "
"Edit verify_environment before spending more sampling budget."
)
continue
consecutive_skips = 0
fwdbwd = await training.forward_backward_async(step_datums, loss_fn="importance_sampling")
optim = await training.optim_step_async(types.AdamParams(learning_rate=LEARNING_RATE))
result = await fwdbwd.result_async()
await optim.result_async()
await checkpoint(training, step)
mean_reward = sum(step_rewards) / max(len(step_rewards), 1)
loss = metric_loss(result)
append_metrics(
args.log,
{"step": step, "mean_reward": mean_reward, "n_datums": len(step_datums), "loss": loss},
)
print(f"step {step:03d} R={mean_reward:.3f} datums={len(step_datums)} loss={loss}")
if eval_rows and step % EVAL_EVERY == 0:
eval_mean = await evaluate(eval_rows, training, tokenizer, params)
append_metrics(
args.log,
{"step": step, "mean_reward": eval_mean, "n_datums": len(eval_rows), "loss": None, "split": "eval"},
)
print(f"step {step:03d} eval_R={eval_mean:.3f} n={len(eval_rows)}")
sampling = await sampling_client(training)
preview_prompt = types.ModelInput.from_ints(tokenizer.encode(preview_prompt_text))
preview = await sampling.sample_async(prompt=preview_prompt, num_samples=1, sampling_params=params)
print("sample:", tokenizer.decode(preview.sequences[0].tokens))
if __name__ == "__main__":
args = parse_args()
apply_args(args)
asyncio.run(train(args))
Loading the builder…