Multi-turn RL: guess a number from 1 to 20
"# Multi-turn RL: guess a number from 1 to 20\n\nThis tutorial builds a small multi-turn environment inspired by PipelineRL's\nguessing task. The model must guess a hidden number in `[1, 20]`.\n\nLoop per rollout:\n1. Model emits a guess as `<answer>N</answer>`.\n2. Environment returns `<feedback>higher</feedback>`, `<feedback>lower</feedback>`,\n or success.\n3. Repeat for a fixed turn budget.\n\nWe train with:\n- `custom_generate_function`: runs the interaction loop.\n- `custom_rm_function`: rewards correct answers with an early-turn bonus.\n- `loss_mask`: trains only on model-generated tokens, not environment feedback.\n\n```python\nimport json\nimport re\n\nfrom modal_training_gym import (\n DatasetConfig,\n Endpoint,\n Qwen3_5_4B,\n Qwen3_5_4B_Recipe,\n TrainConfig,\n)\n```\n\n## Build a deterministic guessing dataset\n\nKeep it simple:\n- train on odd targets\n- evaluate on even targets\n\n```python\n_MAX_VALUE = 20\n_MAX_TURNS = 6\n_PROMPT = (\n \"You are playing a number guessing game.\\n\"\n \"The hidden integer is between 1 and 20.\\n\"\n \"Return only guesses in this exact format: <answer>N</answer> \"\n \"where N is an integer between 1 and 20.\\n\"\n \"After each guess, you will receive <feedback>higher</feedback> or \"\n \"<feedback>lower</feedback>, and must update your next guess accordingly.\"\n)\n\nTRAIN_TARGETS = list(range(1, _MAX_VALUE + 1, 2))\nTEST_TARGETS = list(range(2, _MAX_VALUE + 1, 2))\n\nclass NumberGuessDataset(DatasetConfig):\n input_key = \"messages\"\n label_key = \"label\"\n apply_chat_template = True\n input_column = \"prompt\"\n always_prepare = True # For the purpose of this tutorial, we want to prepare the dataset every time we run it, in case there is stale data from a previous run.\n\n def load(self, split=\"all\"):\n targets = TRAIN_TARGETS if split == \"train\" else TEST_TARGETS\n return [{\"prompt\": _PROMPT, \"target\": target} for target in targets]\n\n def prepare(self, path: str, eval_paths: dict[str, str] | None = None):\n import os\n\n from datasets import Dataset\n\n os.makedirs(os.path.dirname(path), exist_ok=True)\n\n def _row(target: int) -> dict:\n return {\n \"messages\": [{\"role\": \"user\", \"content\": _PROMPT}],\n \"label\": json.dumps({\"answer\": target}),\n }\n\n train_rows = [_row(target) for target in TRAIN_TARGETS for _ in range(20)]\n eval_rows = [_row(target) for target in TEST_TARGETS]\n\n Dataset.from_list(train_rows).to_parquet(path)\n if eval_paths:\n for eval_path in eval_paths.values():\n os.makedirs(os.path.dirname(eval_path), exist_ok=True)\n Dataset.from_list(eval_rows).to_parquet(eval_path)\n\ntrain_dataset = NumberGuessDataset()\n\neval_dataset = NumberGuessDataset()\n```\n\n## Multi-turn environment and reward\n\n`number_guess_generate` is the environment loop:\n- model generates `<answer>N</answer>`\n- environment appends `<feedback>higher|lower</feedback>` when incorrect\n- only model text is trained (`loss_mask=1`), feedback is masked out (`loss_mask=0`)\n\nReward mirrors PipelineRL-style shaping:\n- success: `2.0 - 0.1 * (turns - 1)`\n- malformed output: `-2.0`\n- otherwise: `-1.0`\n\n```python\n_ANSWER_RE = re.compile(r\"<answer>\\s*(\\d+)\\s*</answer>\", re.IGNORECASE)\n\ndef _parse_label(sample) -> dict:\n raw = getattr(sample, \"label\", None)\n if isinstance(raw, dict):\n return raw\n if isinstance(raw, str):\n try:\n return json.loads(raw)\n except json.JSONDecodeError:\n return {}\n return {}\n\ndef _extract_answer(text: str) -> int | None:\n matches = list(_ANSWER_RE.finditer(text))\n if not matches:\n return None\n guess = int(matches[-1].group(1))\n if 1 <= guess <= _MAX_VALUE:\n return guess\n return None\n\nasync def number_guess_generate(args, sample, sampling_params):\n from slime.rollout.sglang_rollout import GenerateState\n from slime.utils.http_utils import post\n from slime.utils.types import Sample\n\n label = _parse_label(sample)\n target = int(label.get(\"answer\", 1))\n max_turns = int(getattr(args, \"max_turns\", _MAX_TURNS))\n\n state = GenerateState(args)\n url = f\"http://{args.sglang_router_ip}:{args.sglang_router_port}/generate\"\n\n prompt_ids = state.tokenizer(sample.prompt, add_special_tokens=False)[\"input_ids\"]\n trajectory_text = \"\"\n response_segments: list[tuple[str, int]] = []\n\n success = False\n format_error = False\n turns_taken = max_turns\n final_status = Sample.Status.COMPLETED\n\n for turn in range(max_turns):\n output = await post(\n url,\n {\n \"text\": f\"{sample.prompt}\\n{trajectory_text}\".strip(),\n \"sampling_params\": sampling_params,\n },\n )\n finish_type = output[\"meta_info\"][\"finish_reason\"][\"type\"]\n if finish_type == \"abort\":\n sample.status = Sample.Status.ABORTED\n return sample\n\n model_text = output[\"text\"]\n trajectory_text += model_text\n response_segments.append((model_text, 1))\n\n guess = _extract_answer(model_text)\n if guess is None:\n format_error = True\n turns_taken = turn + 1\n break\n\n if guess == target:\n success = True\n turns_taken = turn + 1\n break\n\n feedback = \"higher\" if guess < target else \"lower\"\n feedback_text = f\"\\n<feedback>{feedback}</feedback>\\n\"\n trajectory_text += feedback_text\n response_segments.append((feedback_text, 0))\n\n if finish_type == \"length\":\n final_status = Sample.Status.TRUNCATED\n break\n\n response_token_ids: list[int] = []\n loss_masks: list[int] = []\n for segment_text, trainable in response_segments:\n token_ids = state.tokenizer(\n segment_text,\n add_special_tokens=False,\n )[\"input_ids\"]\n response_token_ids += token_ids\n loss_masks += [trainable] * len(token_ids)\n\n sample.tokens = prompt_ids + response_token_ids\n sample.response_length = len(response_token_ids)\n sample.response = trajectory_text\n sample.loss_mask = loss_masks\n sample.status = final_status\n\n sample_metadata = getattr(sample, \"metadata\", None)\n if not isinstance(sample_metadata, dict):\n sample_metadata = {}\n sample_metadata[\"guessing\"] = {\n \"target\": target,\n \"success\": success,\n \"format_error\": format_error,\n \"turns_taken\": turns_taken,\n }\n sample.metadata = sample_metadata\n return sample\n\ndef _trajectory_reward(success: bool, format_error: bool, turns_taken: int) -> float:\n if success:\n return float(2.0 - 0.1 * max(0, turns_taken - 1))\n if format_error:\n return -2.0\n return -1.0\n\nasync def number_guess_rm(args, sample, **kwargs) -> float:\n sample_metadata = getattr(sample, \"metadata\", None)\n guessing_meta = sample_metadata.get(\"guessing\", {}) if isinstance(sample_metadata, dict) else {}\n\n success = bool(guessing_meta.get(\"success\", False))\n format_error = bool(guessing_meta.get(\"format_error\", False))\n turns_taken = int(guessing_meta.get(\"turns_taken\", getattr(args, \"max_turns\", _MAX_TURNS)))\n return _trajectory_reward(\n success=success,\n format_error=format_error,\n turns_taken=turns_taken,\n )\n```\n\n## Offline multi-turn trajectory evaluator\n\nWe create a full multi-turn evaluator that aggregates scores in a\nsmall loop over the eval dataset.\n\n```python\ndef run_guessing_trajectory(\n deployment: Endpoint,\n *,\n target: int,\n max_turns: int = _MAX_TURNS,\n) -> dict:\n trace = \"\"\n for turn in range(max_turns):\n prompt = f\"{_PROMPT}\\n{trace}\".strip()\n msg = deployment.chat(\n [{\"role\": \"user\", \"content\": prompt}],\n chat_template_kwargs={\"enable_thinking\": False},\n )\n response = msg.get(\"content\") or msg.get(\"reasoning_content\") or \"\"\n guess = _extract_answer(response)\n if guess is None:\n return {\n \"success\": False,\n \"format_error\": True,\n \"turns_taken\": turn + 1,\n \"response\": response,\n }\n if guess == target:\n return {\n \"success\": True,\n \"format_error\": False,\n \"turns_taken\": turn + 1,\n \"response\": response,\n }\n feedback = \"higher\" if guess < target else \"lower\"\n trace += f\"{response}\\n<feedback>{feedback}</feedback>\\n\"\n return {\n \"success\": False,\n \"format_error\": False,\n \"turns_taken\": max_turns,\n \"response\": trace,\n }\n\ndef guessing_eval_fn(\n deployment: Endpoint,\n example: dict,\n) -> dict:\n target = int(example[\"target\"])\n trajectory = run_guessing_trajectory(\n deployment,\n target=target,\n max_turns=_MAX_TURNS,\n )\n reward = _trajectory_reward(\n success=trajectory[\"success\"],\n format_error=trajectory[\"format_error\"],\n turns_taken=trajectory[\"turns_taken\"],\n )\n return {\n \"score\": reward,\n \"response\": trajectory[\"response\"],\n \"metadata\": {\n \"success\": trajectory[\"success\"],\n \"format_error\": trajectory[\"format_error\"],\n \"turns_taken\": trajectory[\"turns_taken\"],\n \"target\": target,\n },\n }\n\ndef summarize_eval(rows: list[dict]) -> dict:\n success_rate = sum(1 for row in rows if row[\"metadata\"].get(\"success\")) / max(\n len(rows), 1\n )\n mean_turns = sum(\n int(row[\"metadata\"].get(\"turns_taken\", _MAX_TURNS)) for row in rows\n ) / max(len(rows), 1)\n return {\n \"success_rate\": float(success_rate),\n \"mean_turns\": float(mean_turns),\n }\n\ndef run_eval(\n deployment, *, max_concurrency: int = 2\n) -> tuple[float, list[dict]]:\n from concurrent.futures import ThreadPoolExecutor\n\n deployment.wait_until_ready(timeout=15 * 60)\n\n def _score_one(example):\n return guessing_eval_fn(deployment, example)\n\n with ThreadPoolExecutor(max_workers=max_concurrency) as executor:\n rows = list(executor.map(_score_one, eval_dataset.load()))\n mean = sum(r[\"score\"] for r in rows) / len(rows) if rows else float(\"nan\")\n return mean, rows\n```\n\n## Serve and evaluate the base model\n\n```python\nmodel = Qwen3_5_4B()\nbase_deployment = Endpoint.launch(\n model,\n unauthenticated=True,\n recreate_if_existing=True,\n)\nprint(f\"Base model URL: {base_deployment.url}\")\nbase_mean, base_rows = run_eval(base_deployment)\nbase_summary = summarize_eval(base_rows)\nprint(f\"Base success rate: {base_summary['success_rate']:.2%}\")\nprint(f\"Base mean reward: {base_mean:.3f}\")\nprint(f\"Base mean turns: {base_summary['mean_turns']:.2f}\")\n```\n\n## Train with custom multi-turn rollout\n\nA quick tour of the knobs we set below.\n\n**Cluster and parallelism**\n- `gpu_type=\"H100\"` — GPU SKU used for both the rollout (sglang) and training\n (Megatron) ranks.\n- `colocate=True` — share the same GPUs between rollout and training, alternating\n between the two. Set `False` to give sglang dedicated GPUs (faster, more expensive).\n- `tensor_model_parallel_size=1` — Megatron tensor-parallel degree. The 4B\n preset still uses 8 H100s; bump TP for larger models that outgrow one GPU.\n- `sequence_parallel=False` — only meaningful when `tensor_model_parallel_size > 1`.\n- `rollout_num_gpus_per_engine=1` — GPUs per sglang inference engine (sglang's TP).\n\n**Rollout**\n- `num_rollout=20` — total rollout/train iterations to run. Each iteration samples\n a batch, scores it, and applies one policy update.\n- `rollout_batch_size=8` — prompts sampled per rollout iteration.\n- `rollout_max_response_len=64` — max new tokens per sglang call. We keep it tiny\n because every turn is `<answer>N</answer>` plus a bit of thinking.\n- `rollout_temperature=1.0` — sampling temperature during rollouts.\n\n**Training and checkpoints**\n- `global_batch_size=8` — effective batch size for the policy gradient update.\n- `save_interval=10` — write a Megatron checkpoint every N rollout iterations.\n- `apply_chat_template_kwargs='{\"enable_thinking\": false}'` — passed to the\n tokenizer's chat template; disables Qwen3's `<think>` block so responses stay\n short and parseable.\n\n```python\nconfig = TrainConfig(\n model=model,\n dataset=train_dataset,\n recipe=Qwen3_5_4B_Recipe(\n eval_interval=None,\n custom_generate_function=number_guess_generate,\n custom_rm_function=number_guess_rm,\n extra_config={\n \"max_turns\": _MAX_TURNS,\n \"log_multi_turn\": True,\n },\n\n gpu_type=\"H100\",\n colocate=True,\n tensor_model_parallel_size=1,\n sequence_parallel=False,\n rollout_num_gpus_per_engine=1,\n\n num_rollout=20,\n rollout_batch_size=8,\n n_samples_per_prompt=4,\n rollout_max_response_len=64,\n rollout_temperature=1.0,\n\n global_batch_size=8,\n save_interval=10,\n apply_chat_template_kwargs='{\"enable_thinking\": false}',\n ),\n)\nprint(\"Starting training...\")\nrun = config.launch()\nprint(f\"run id: {run.training_run_id}\")\n```\n\n## Evaluate trained checkpoint\n\n```python\nresult = run.result()\ncheckpoint = result.checkpoints()[-1]\ntrained_deployment = Endpoint.launch(\n model, checkpoint, unauthenticated=True, recreate_if_existing=True\n)\nprint(f\"Trained model URL: {trained_deployment.url}\")\n\ntrained_mean, trained_rows = run_eval(trained_deployment)\ntrained_summary = summarize_eval(trained_rows)\nprint(f\"Trained success rate: {trained_summary['success_rate']:.2%}\")\nprint(f\"Trained mean reward: {trained_mean:.3f}\")\nprint(f\"Trained mean turns: {trained_summary['mean_turns']:.2f}\")\nprint(f\"Base success rate: {base_summary['success_rate']:.2%}\")\nprint(f\"Base mean reward: {base_mean:.3f}\")\nprint(f\"Base mean turns: {base_summary['mean_turns']:.2f}\")\n```\n"
This tutorial builds a small multi-turn environment inspired by PipelineRL’s
guessing task. The model must guess a hidden number in [1, 20].
Loop per rollout:
- Model emits a guess as
<answer>N</answer>. - Environment returns
<feedback>higher</feedback>,<feedback>lower</feedback>, or success. - Repeat for a fixed turn budget.
We train with:
custom_generate_function: runs the interaction loop.custom_rm_function: rewards correct answers with an early-turn bonus.loss_mask: trains only on model-generated tokens, not environment feedback.
import jsonimport re
from modal_training_gym import ( DatasetConfig, Endpoint, Qwen3_5_4B, Qwen3_5_4B_Recipe, TrainConfig,)Build a deterministic guessing dataset
Keep it simple:
- train on odd targets
- evaluate on even targets
_MAX_VALUE = 20_MAX_TURNS = 6_PROMPT = ( "You are playing a number guessing game.\n" "The hidden integer is between 1 and 20.\n" "Return only guesses in this exact format: <answer>N</answer> " "where N is an integer between 1 and 20.\n" "After each guess, you will receive <feedback>higher</feedback> or " "<feedback>lower</feedback>, and must update your next guess accordingly.")
TRAIN_TARGETS = list(range(1, _MAX_VALUE + 1, 2))TEST_TARGETS = list(range(2, _MAX_VALUE + 1, 2))
class NumberGuessDataset(DatasetConfig): input_key = "messages" label_key = "label" apply_chat_template = True input_column = "prompt" always_prepare = True # For the purpose of this tutorial, we want to prepare the dataset every time we run it, in case there is stale data from a previous run.
def load(self, split="all"): targets = TRAIN_TARGETS if split == "train" else TEST_TARGETS return [{"prompt": _PROMPT, "target": target} for target in targets]
def prepare(self, path: str, eval_paths: dict[str, str] | None = None): import os
from datasets import Dataset
os.makedirs(os.path.dirname(path), exist_ok=True)
def _row(target: int) -> dict: return { "messages": [{"role": "user", "content": _PROMPT}], "label": json.dumps({"answer": target}), }
train_rows = [_row(target) for target in TRAIN_TARGETS for _ in range(20)] eval_rows = [_row(target) for target in TEST_TARGETS]
Dataset.from_list(train_rows).to_parquet(path) if eval_paths: for eval_path in eval_paths.values(): os.makedirs(os.path.dirname(eval_path), exist_ok=True) Dataset.from_list(eval_rows).to_parquet(eval_path)
train_dataset = NumberGuessDataset()
eval_dataset = NumberGuessDataset()Multi-turn environment and reward
number_guess_generate is the environment loop:
- model generates
<answer>N</answer> - environment appends
<feedback>higher|lower</feedback>when incorrect - only model text is trained (
loss_mask=1), feedback is masked out (loss_mask=0)
Reward mirrors PipelineRL-style shaping:
- success:
2.0 - 0.1 * (turns - 1) - malformed output:
-2.0 - otherwise:
-1.0
_ANSWER_RE = re.compile(r"<answer>\s*(\d+)\s*</answer>", re.IGNORECASE)
def _parse_label(sample) -> dict: raw = getattr(sample, "label", None) if isinstance(raw, dict): return raw if isinstance(raw, str): try: return json.loads(raw) except json.JSONDecodeError: return {} return {}
def _extract_answer(text: str) -> int | None: matches = list(_ANSWER_RE.finditer(text)) if not matches: return None guess = int(matches[-1].group(1)) if 1 <= guess <= _MAX_VALUE: return guess return None
async def number_guess_generate(args, sample, sampling_params): from slime.rollout.sglang_rollout import GenerateState from slime.utils.http_utils import post from slime.utils.types import Sample
label = _parse_label(sample) target = int(label.get("answer", 1)) max_turns = int(getattr(args, "max_turns", _MAX_TURNS))
state = GenerateState(args) url = f"http://{args.sglang_router_ip}:{args.sglang_router_port}/generate"
prompt_ids = state.tokenizer(sample.prompt, add_special_tokens=False)["input_ids"] trajectory_text = "" response_segments: list[tuple[str, int]] = []
success = False format_error = False turns_taken = max_turns final_status = Sample.Status.COMPLETED
for turn in range(max_turns): output = await post( url, { "text": f"{sample.prompt}\n{trajectory_text}".strip(), "sampling_params": sampling_params, }, ) finish_type = output["meta_info"]["finish_reason"]["type"] if finish_type == "abort": sample.status = Sample.Status.ABORTED return sample
model_text = output["text"] trajectory_text += model_text response_segments.append((model_text, 1))
guess = _extract_answer(model_text) if guess is None: format_error = True turns_taken = turn + 1 break
if guess == target: success = True turns_taken = turn + 1 break
feedback = "higher" if guess < target else "lower" feedback_text = f"\n<feedback>{feedback}</feedback>\n" trajectory_text += feedback_text response_segments.append((feedback_text, 0))
if finish_type == "length": final_status = Sample.Status.TRUNCATED break
response_token_ids: list[int] = [] loss_masks: list[int] = [] for segment_text, trainable in response_segments: token_ids = state.tokenizer( segment_text, add_special_tokens=False, )["input_ids"] response_token_ids += token_ids loss_masks += [trainable] * len(token_ids)
sample.tokens = prompt_ids + response_token_ids sample.response_length = len(response_token_ids) sample.response = trajectory_text sample.loss_mask = loss_masks sample.status = final_status
sample_metadata = getattr(sample, "metadata", None) if not isinstance(sample_metadata, dict): sample_metadata = {} sample_metadata["guessing"] = { "target": target, "success": success, "format_error": format_error, "turns_taken": turns_taken, } sample.metadata = sample_metadata return sample
def _trajectory_reward(success: bool, format_error: bool, turns_taken: int) -> float: if success: return float(2.0 - 0.1 * max(0, turns_taken - 1)) if format_error: return -2.0 return -1.0
async def number_guess_rm(args, sample, **kwargs) -> float: sample_metadata = getattr(sample, "metadata", None) guessing_meta = sample_metadata.get("guessing", {}) if isinstance(sample_metadata, dict) else {}
success = bool(guessing_meta.get("success", False)) format_error = bool(guessing_meta.get("format_error", False)) turns_taken = int(guessing_meta.get("turns_taken", getattr(args, "max_turns", _MAX_TURNS))) return _trajectory_reward( success=success, format_error=format_error, turns_taken=turns_taken, )Offline multi-turn trajectory evaluator
We create a full multi-turn evaluator that aggregates scores in a small loop over the eval dataset.
def run_guessing_trajectory( deployment: Endpoint, *, target: int, max_turns: int = _MAX_TURNS,) -> dict: trace = "" for turn in range(max_turns): prompt = f"{_PROMPT}\n{trace}".strip() msg = deployment.chat( [{"role": "user", "content": prompt}], chat_template_kwargs={"enable_thinking": False}, ) response = msg.get("content") or msg.get("reasoning_content") or "" guess = _extract_answer(response) if guess is None: return { "success": False, "format_error": True, "turns_taken": turn + 1, "response": response, } if guess == target: return { "success": True, "format_error": False, "turns_taken": turn + 1, "response": response, } feedback = "higher" if guess < target else "lower" trace += f"{response}\n<feedback>{feedback}</feedback>\n" return { "success": False, "format_error": False, "turns_taken": max_turns, "response": trace, }
def guessing_eval_fn( deployment: Endpoint, example: dict,) -> dict: target = int(example["target"]) trajectory = run_guessing_trajectory( deployment, target=target, max_turns=_MAX_TURNS, ) reward = _trajectory_reward( success=trajectory["success"], format_error=trajectory["format_error"], turns_taken=trajectory["turns_taken"], ) return { "score": reward, "response": trajectory["response"], "metadata": { "success": trajectory["success"], "format_error": trajectory["format_error"], "turns_taken": trajectory["turns_taken"], "target": target, }, }
def summarize_eval(rows: list[dict]) -> dict: success_rate = sum(1 for row in rows if row["metadata"].get("success")) / max( len(rows), 1 ) mean_turns = sum( int(row["metadata"].get("turns_taken", _MAX_TURNS)) for row in rows ) / max(len(rows), 1) return { "success_rate": float(success_rate), "mean_turns": float(mean_turns), }
def run_eval( deployment, *, max_concurrency: int = 2) -> tuple[float, list[dict]]: from concurrent.futures import ThreadPoolExecutor
deployment.wait_until_ready(timeout=15 * 60)
def _score_one(example): return guessing_eval_fn(deployment, example)
with ThreadPoolExecutor(max_workers=max_concurrency) as executor: rows = list(executor.map(_score_one, eval_dataset.load())) mean = sum(r["score"] for r in rows) / len(rows) if rows else float("nan") return mean, rowsServe and evaluate the base model
model = Qwen3_5_4B()base_deployment = Endpoint.launch( model, unauthenticated=True, recreate_if_existing=True,)print(f"Base model URL: {base_deployment.url}")base_mean, base_rows = run_eval(base_deployment)base_summary = summarize_eval(base_rows)print(f"Base success rate: {base_summary['success_rate']:.2%}")print(f"Base mean reward: {base_mean:.3f}")print(f"Base mean turns: {base_summary['mean_turns']:.2f}")Train with custom multi-turn rollout
A quick tour of the knobs we set below.
Cluster and parallelism
gpu_type="H100"— GPU SKU used for both the rollout (sglang) and training (Megatron) ranks.colocate=True— share the same GPUs between rollout and training, alternating between the two. SetFalseto give sglang dedicated GPUs (faster, more expensive).tensor_model_parallel_size=1— Megatron tensor-parallel degree. The 4B preset still uses 8 H100s; bump TP for larger models that outgrow one GPU.sequence_parallel=False— only meaningful whentensor_model_parallel_size > 1.rollout_num_gpus_per_engine=1— GPUs per sglang inference engine (sglang’s TP).
Rollout
num_rollout=20— total rollout/train iterations to run. Each iteration samples a batch, scores it, and applies one policy update.rollout_batch_size=8— prompts sampled per rollout iteration.rollout_max_response_len=64— max new tokens per sglang call. We keep it tiny because every turn is<answer>N</answer>plus a bit of thinking.rollout_temperature=1.0— sampling temperature during rollouts.
Training and checkpoints
global_batch_size=8— effective batch size for the policy gradient update.save_interval=10— write a Megatron checkpoint every N rollout iterations.apply_chat_template_kwargs='{"enable_thinking": false}'— passed to the tokenizer’s chat template; disables Qwen3’s<think>block so responses stay short and parseable.
config = TrainConfig( model=model, dataset=train_dataset, recipe=Qwen3_5_4B_Recipe( eval_interval=None, custom_generate_function=number_guess_generate, custom_rm_function=number_guess_rm, extra_config={ "max_turns": _MAX_TURNS, "log_multi_turn": True, },
gpu_type="H100", colocate=True, tensor_model_parallel_size=1, sequence_parallel=False, rollout_num_gpus_per_engine=1,
num_rollout=20, rollout_batch_size=8, n_samples_per_prompt=4, rollout_max_response_len=64, rollout_temperature=1.0,
global_batch_size=8, save_interval=10, apply_chat_template_kwargs='{"enable_thinking": false}', ),)print("Starting training...")run = config.launch()print(f"run id: {run.training_run_id}")Evaluate trained checkpoint
result = run.result()checkpoint = result.checkpoints()[-1]trained_deployment = Endpoint.launch( model, checkpoint, unauthenticated=True, recreate_if_existing=True)print(f"Trained model URL: {trained_deployment.url}")
trained_mean, trained_rows = run_eval(trained_deployment)trained_summary = summarize_eval(trained_rows)print(f"Trained success rate: {trained_summary['success_rate']:.2%}")print(f"Trained mean reward: {trained_mean:.3f}")print(f"Trained mean turns: {trained_summary['mean_turns']:.2f}")print(f"Base success rate: {base_summary['success_rate']:.2%}")print(f"Base mean reward: {base_mean:.3f}")print(f"Base mean turns: {base_summary['mean_turns']:.2f}")