gpt-5.6-sol
on-policy-self-distillation — graded by the task's verifier against a measured baseline.
- backend
- modal
- variant
- codex azure_ai/gpt-5.6-sol
- metric
- AIME24 avg@12 accuracy
- actions
- 636
- duration
- 8h 10m
- recorded
- 2026-08-03
run idopsd-aime24-4gpus__3tyLi4e
Everything gpt-5.6-sol did on this task, in order.
download JSON ↗encrypted chain-of-thought — recorded, not exposed
args
{
"input": "const r = await tools.exec_command({\n cmd: \"find /opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd -type f -newermt '2026-08-03 00:00:00' -printf '%p\\\\n' | sort\\nfind /opt/hf_cache/modules -type f -printf '%p\\\\n' 2>/dev/null | sort\\nfind /opt/hf_cache/datasets -maxdepth 1 -type f -newermt '2026-08-03 00:00:00' -printf '%p\\\\n' | sort\\nls -lt /workspace/eval_results | head -20\\ncat /workspace/.timer/remaining_secs\",\n workdir: \"/workspace/repo\",\n yield_time_ms: 10000,\n max_output_tokens: 6000\n});\ntext(r.output);"
}exec result
Script completed
Wall time 0.2 seconds
Output:
/opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/cache-55f01c4829935f27.arrow
/opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/cache-a47d60ea3fed15a4.arrow
/opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/cache-d162f0b4cc47513f.arrow
/opt/hf_cache/modules/__init__.py
/opt/hf_cache/modules/datasets_modules/__init__.py
/opt/hf_cache/datasets/_opt_hf_cache_datasets_siyanzhao___openthoughts_math_30k_opsd_default_0.0.0_1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62.lock
total 259
-rw-r--r-- 1 root root 37531 Aug 3 09:24 dev_eval_1785747779.json
-rw-r--r-- 1 root root 37497 Aug 3 08:43 frontier_v4_clip002_full.json
-rw-r--r-- 1 root root 37505 Aug 3 08:02 frontier_v3_aops_full.json
-rw-r--r-- 1 root root 37519 Aug 3 07:20 frontier_v2_rep2_full.json
-rw-r--r-- 1 root root 37466 Aug 3 06:41 frontier_v2_full.json
-rw-r--r-- 1 root root 37474 Aug 3 06:01 frontier_v1_rep2_full.json
-rw-r--r-- 1 root root 37519 Aug 3 05:23 frontier_v1_full.json
4398
encrypted chain-of-thought — recorded, not exposed
args
{
"input": "const r = await tools.exec_command({\n cmd: \"unlink /opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/cache-55f01c4829935f27.arrow\\nunlink /opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/cache-a47d60ea3fed15a4.arrow\\nunlink /opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/cache-d162f0b4cc47513f.arrow\\nunlink /opt/hf_cache/datasets/_opt_hf_cache_datasets_siyanzhao___openthoughts_math_30k_opsd_default_0.0.0_1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62.lock\\nunlink /opt/hf_cache/modules/datasets_modules/__init__.py\\nunlink /opt/hf_cache/modules/__init__.py\\nrmdir /opt/hf_cache/modules/datasets_modules /opt/hf_cache/modules\\njq '{accuracy, correct, total, format_rate, checkpoint}' /workspace/eval_results/frontier_v1_full.json /workspace/eval_results/frontier_v1_rep2_full.json /workspace/eval_results/dev_eval_1785747779.json\\nsha256sum /opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/openthoughts_math_30k_opsd-train-00000-of-00002.arrow /opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/openthoughts_math_30k_opsd-train-00001-of-00002.arrow /opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/dataset_info.json\",\n workdir: \"/workspace/repo\",\n yield_time_ms: 10000,\n max_output_tokens: 4000\n});\ntext(r.output);"
}exec result
Script completed
Wall time 3.4 seconds
Output:
{
"accuracy": null,
"correct": null,
"total": null,
"format_rate": null,
"checkpoint": null
}
{
"accuracy": null,
"correct": null,
"total": null,
"format_rate": null,
"checkpoint": null
}
{
"accuracy": null,
"correct": null,
"total": null,
"format_rate": null,
"checkpoint": null
}
da57b57c57b35fe5d98a61273603ce1a7513a42368b6ed209c22409338369ead /opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/openthoughts_math_30k_opsd-train-00000-of-00002.arrow
727ba2b8c54d190651709e5a5f323e844fc90dcbf03309197950ef070595f2b5 /opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/openthoughts_math_30k_opsd-train-00001-of-00002.arrow
1535163387e4c5da3a21aa51766d31261806e335d1cb141104cb285aa7b3d76c /opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd/default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/dataset_info.json
encrypted chain-of-thought — recorded, not exposed
args
{
"input": "const r = await tools.exec_command({\n cmd: \"jq 'keys' /workspace/eval_results/frontier_v1_full.json\\njq '.summary // .metrics // .results // .[0]' /workspace/eval_results/frontier_v1_full.json | head -100\\npython -m py_compile opsd_train.py opsd_trainer.py data_collator.py\\ngit diff --check -- opsd_train.py opsd_trainer.py data_collator.py\\ngit status --short\\nfind /opt/hf_cache/datasets/siyanzhao___openthoughts_math_30k_opsd -type f -printf '%P %s %TY-%Tm-%TdT%TH:%TM:%TS\\\\n' | sort\\nfind /opt/hf_cache -type f -newermt '2026-08-03 00:00:00' -printf '%p\\\\n' | sort | head -100\\ncat /workspace/submission/recipe.env\",\n workdir: \"/workspace/repo\",\n yield_time_ms: 10000,\n max_output_tokens: 6000\n});\ntext(r.output);"
}exec result
Script completed
Wall time 0.3 seconds
Output:
[
"average_at_n_pct",
"benchmark",
"checkpoint_kind",
"checkpoint_path",
"data_parallel_size",
"format_rate_pct",
"frozen_settings",
"metric",
"num_problems",
"pass_at_n_pct",
"results",
"seed",
"total_correct",
"total_solutions"
]
[
{
"problem_id": 60,
"ground_truth": "204",
"num_correct": 12,
"n": 12,
"majority_vote_correct": true,
"generations": [
{
"predicted_answer": "204",
"correct": true,
"formatted": true
},
{
"predicted_answer": "204",
"correct": true,
"formatted": true
},
{
"predicted_answer": "204",
"correct": true,
"formatted": true
},
{
"predicted_answer": "204",
"correct": true,
"formatted": true
},
{
"predicted_answer": "204",
"correct": true,
"formatted": true
},
{
"predicted_answer": "204",
"correct": true,
"formatted": true
},
{
"predicted_answer": "204",
"correct": true,
"formatted": true
},
{
"predicted_answer": "204",
"correct": true,
"formatted": true
},
{
"predicted_answer": "204",
"correct": true,
"formatted": true
},
{
"predicted_answer": "204",
"correct": true,
"formatted": true
},
{
"predicted_answer": "204",
"correct": true,
"formatted": true
},
{
"predicted_answer": "204",
"correct": true,
"formatted": true
}
]
},
{
"problem_id": 61,
"ground_truth": "113",
"num_correct": 10,
"n": 12,
"majority_vote_correct": true,
"generations": [
{
"predicted_answer": "113",
"correct": true,
"formatted": true
},
{
"predicted_answer": "113",
"correct": true,
"formatted": true
},
{
"predicted_answer": "119",
"correct": false,
"formatted": true
},
{
"predicted_answer": "113",
"correct": true,
"formatted": true
},
{
"predicted_answer": "113",
"correct": true,
M README.md
M data_collator.py
M opsd_train.py
M opsd_trainer.py
default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62.incomplete_info.lock 0 2026-08-02T17:14:36.7049412940
default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/dataset_info.json 1562 2026-08-02T17:14:36.7039174830
default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/openthoughts_math_30k_opsd-train-00000-of-00002.arrow 537225960 2026-08-02T17:14:36.6819176270
default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62/openthoughts_math_30k_opsd-train-00001-of-00002.arrow 116664552 2026-08-02T17:14:36.7019174960
default/0.0.0/1f33e9dc2e8a1c639ca74f8024ad4a9f1f5eae62_builder.lock 0 2026-08-02T17:14:28.6370304490
lmbda=0.01
lr_scheduler_type=linear
encrypted chain-of-thought — recorded, not exposed
args
{
"input": "const r = await tools.exec_command({\n cmd: \"for f in /workspace/eval_results/frontier_v1_full.json /workspace/eval_results/frontier_v1_rep2_full.json /workspace/eval_results/dev_eval_1785747779.json; do jq -r '[input_filename, .average_at_n_pct, .total_correct, .total_solutions, .format_rate_pct] | @tsv' \\\"$f\\\"; done\\nawk 'BEGIN { print (56.3888888889 + 53.8888888889 + 54.7222222222) / 3; print ((56.3888888889 + 53.8888888889 + 54.7222222222) / 3) - 54.22 }'\\ngit diff --stat -- opsd_train.py opsd_trainer.py data_collator.py\\ncat /workspace/.timer/remaining_secs\",\n workdir: \"/workspace/repo\",\n yield_time_ms: 10000,\n max_output_tokens: 3000\n});\ntext(r.output);"
}exec result
Script completed
Wall time 0.2 seconds
Output:
/workspace/eval_results/frontier_v1_full.json 56.388888888888886 203 360 99.72222222222223
/workspace/eval_results/frontier_v1_rep2_full.json 53.888888888888886 194 360 100
/workspace/eval_results/dev_eval_1785747779.json 54.72222222222222 197 360 99.72222222222223
55
0.78
data_collator.py | 4 +++
opsd_train.py | 21 ++++++++++++++
opsd_trainer.py | 88 ++++++++++++++++++++++++++++++++++++++++++++++++++++----
3 files changed, 107 insertions(+), 6 deletions(-)
4368
The three-run estimate confirms the improvement is real but modest. Before freezing the submission, I’m checking one subtle implementation risk—the correspondence between vLLM completions and reference answers across ranks—because a misalignment there would turn the correctness gate into noise even if aggregate training still looks healthy.
encrypted chain-of-thought — recorded, not exposed
args
{
"input": "const r = await tools.exec_command({\n cmd: \"rg -n \\\"def _generate_on_policy_outputs_vllm|completion_texts|prompt_texts|gather_object|broadcast_object_list\\\" opsd_trainer.py | head -80\\nsed -n '1030,1270p' opsd_trainer.py\",\n workdir: \"/workspace/repo\",\n yield_time_ms: 10000,\n max_output_tokens: 16000\n});\ntext(r.output);"
}exec result
Script completed
Wall time 0.1 seconds
Output:
30:from accelerate.utils import DistributedType, broadcast_object_list, gather_object, is_peft_model
920: def _generate_on_policy_outputs_vllm(self, inputs, generation_config, pad_token_id=None):
961: all_prompts_text = gather_object(prompts_text_for_vllm)
977: completion_ids = broadcast_object_list(completion_ids, from_process=0)
1007: torch.distributed.all_gather_object(
1101: completion_texts = []
1104: completion_texts.append(completion_text)
1106: return new_input_ids, new_attention_mask, new_labels, prompts_text_with_special, completion_texts
1136: all_prompts_text = gather_object(prompts_text)
1148: completion_ids = broadcast_object_list(completion_ids, from_process=0)
1167: torch.distributed.all_gather_object(gathered_prompts, prompts_text, group=self.vllm_tp_group)
1435: generated_ids, generated_attention_mask, _, prompt_texts, completion_texts = result
1443: prompt_texts = self.processing_class.batch_decode(
1448: completion_texts = self.processing_class.batch_decode(
1492: for completion, answer in zip(completion_texts, inputs["answer_texts"])
1498: self._textual_logs["prompt"].extend(gather_object(prompt_texts))
1499: self._textual_logs["completion"].extend(gather_object(completion_texts))
1502: for prompt, completion in zip(prompt_texts, completion_texts):
1512: sample_idx = random.randint(0, len(prompt_texts) - 1)
1513: print(f"\nPrompt:\n{prompt_texts[sample_idx]}")
1514: print(f"\nCompletion:\n{completion_texts[sample_idx]}")
# Calculate and print vLLM generation statistics
elapsed_time = time.time() - start_time
total_completion_tokens = sum(len(ids) for ids in completion_ids)
num_prompts = len(completion_ids)
avg_completion_length = total_completion_tokens / num_prompts if num_prompts > 0 else 0
tokens_per_sec = total_completion_tokens / elapsed_time if elapsed_time > 0 else 0
print(
f"vLLM generation done - elapsed time: {elapsed_time:.2f}s, prompts: {num_prompts}, total tokens: {total_completion_tokens}, avg length: {avg_completion_length:.1f}, speed: {tokens_per_sec:.1f} tok/s"
)
# We need to combine prompt and completion for new_input_ids
# Tokenize prompts again to get prompt_ids on the correct device and format
# Use prompts_text_for_vllm (without special tokens) for tokenization since vLLM expects clean text
# Ensure add_special_tokens=False as vLLM typically handles prompts as raw text
# Calculate max_length for prompts, ensuring it's positive
prompt_max_length = (
max(1, self.args.max_length - max_completion_length) if self.args.max_length else None
)
prompt_tokenized = self.processing_class(
prompts_text_for_vllm,
return_tensors="pt",
padding="longest",
truncation=True if prompt_max_length else False,
max_length=prompt_max_length,
add_special_tokens=False,
).to(device)
prompt_ids = prompt_tokenized.input_ids
completion_ids_tensors = [torch.tensor(ids, device=device) for ids in completion_ids]
# Manually pad/truncate completions to max_completion_length length before using pad function
padded_completion_ids_list = []
for completion_tensor in completion_ids_tensors:
if len(completion_tensor) > max_completion_length:
# Truncate if longer than max_completion_length
padded_completion_ids_list.append(completion_tensor[:max_completion_length])
elif len(completion_tensor) < max_completion_length:
# Pad if shorter than max_completion_length
padding_needed = max_completion_length - len(completion_tensor)
padded_tensor = torch.cat(
[
completion_tensor,
torch.full(
(padding_needed,), pad_token_id, device=device, dtype=completion_tensor.dtype
),
]
)
padded_completion_ids_list.append(padded_tensor)
else:
# Already the right length
padded_completion_ids_list.append(completion_tensor)
# Now all tensors are the same length, so we can stack them
padded_completion_ids = torch.stack(padded_completion_ids_list)
# Ensure prompt_ids and padded_completion_ids are 2D
if prompt_ids.ndim == 1:
prompt_ids = prompt_ids.unsqueeze(0)
if padded_completion_ids.ndim == 1:
padded_completion_ids = padded_completion_ids.unsqueeze(0)
new_input_ids = torch.cat([prompt_ids, padded_completion_ids], dim=1)
new_attention_mask = torch.ones_like(new_input_ids, device=device)
new_labels = new_input_ids.clone()
if pad_token_id is not None:
new_labels[new_labels == pad_token_id] = -100
new_attention_mask[new_input_ids == pad_token_id] = 0
# Extract completion texts from the generated completion IDs
completion_texts = []
for comp_ids in completion_ids:
completion_text = self.processing_class.decode(comp_ids, skip_special_tokens=False)
completion_texts.append(completion_text)
return new_input_ids, new_attention_mask, new_labels, prompts_text_with_special, completion_texts
def _generate_teacher_reasoning_vllm(
self, teacher_reasoning_prompts, teacher_reasoning_attention_mask=None
):
"""Generate teacher's reasoning using vLLM."""
import time
device = self.accelerator.device
# Decode prompts for vLLM
prompts_text = self.processing_class.batch_decode(
teacher_reasoning_prompts,
skip_special_tokens=True,
)
if self.processing_class.pad_token:
prompts_text = [p.replace(self.processing_class.pad_token, "") for p in prompts_text]
max_reasoning_length = self.reasoning_generation_config.max_new_tokens
temperature = self.reasoning_generation_config.temperature
top_k = (
self.reasoning_generation_config.top_k
if self.reasoning_generation_config.top_k and self.reasoning_generation_config.top_k > 0
else -1
)
top_p = self.args.top_p if hasattr(self.args, "top_p") else 1.0
start_time = time.time()
if self.vllm_mode == "server":
all_prompts_text = gather_object(prompts_text)
if self.accelerator.is_main_process:
completion_ids = self.vllm_client.generate(
prompts=all_prompts_text,
n=1,
temperature=temperature,
top_p=top_p,
top_k=top_k,
max_tokens=max_reasoning_length,
)
else:
completion_ids = [None] * len(all_prompts_text)
completion_ids = broadcast_object_list(completion_ids, from_process=0)
process_slice = slice(
self.accelerator.process_index * len(prompts_text),
(self.accelerator.process_index + 1) * len(prompts_text),
)
completion_ids = completion_ids[process_slice]
elif self.vllm_mode == "colocate":
sampling_params = SamplingParams(
n=1,
temperature=temperature,
top_p=top_p,
top_k=top_k,
max_tokens=max_reasoning_length,
)
if hasattr(self, "vllm_tp_group") and self.vllm_tensor_parallel_size > 1:
orig_size = len(prompts_text)
gathered_prompts = [None for _ in range(self.vllm_tensor_parallel_size)]
torch.distributed.all_gather_object(gathered_prompts, prompts_text, group=self.vllm_tp_group)
all_prompts_text = [p for sublist in gathered_prompts for p in sublist]
else:
all_prompts_text = prompts_text
all_outputs = self.vllm_engine.generate(
all_prompts_text, sampling_params=sampling_params, use_tqdm=False
)
completion_ids = [output.token_ids for outputs in all_outputs for output in outputs.outputs]
if hasattr(self, "vllm_tp_group") and self.vllm_tensor_parallel_size > 1:
local_rank_in_group = torch.distributed.get_rank(group=self.vllm_tp_group)
tp_slice = slice(local_rank_in_group * orig_size, (local_rank_in_group + 1) * orig_size)
completion_ids = completion_ids[tp_slice]
if self.vllm_enable_sleep_mode:
self.vllm_engine.sleep(level=2)
elapsed_time = time.time() - start_time
total_tokens = sum(len(ids) for ids in completion_ids)
num_prompts = len(completion_ids)
print(
f"vLLM teacher reasoning generation done - elapsed: {elapsed_time:.2f}s, prompts: {num_prompts}, tokens: {total_tokens}, speed: {total_tokens/elapsed_time:.1f} tok/s"
)
# Combine prompt + completion
prompt_tokenized = self.processing_class(
prompts_text,
return_tensors="pt",
padding="longest",
truncation=True,
add_special_tokens=False,
).to(device)
prompt_ids = prompt_tokenized.input_ids
completion_ids_tensors = [torch.tensor(ids, device=device) for ids in completion_ids]
padded_completions = pad(
completion_ids_tensors, padding_value=self.processing_class.pad_token_id, padding_side="right"
)
reasoning_ids = torch.cat([prompt_ids, padded_completions], dim=1)
return reasoning_ids
def _sync_fsdp_params_to_vllm(self, module: nn.Module, prefix: str = "", visited=None):
"""Memory-efficient post-order traversal of FSDP modules to extract full parameters and sync with student vLLM."""
if visited is None:
visited = set()
for child_name, child_module in module.named_children():
child_prefix = f"{prefix}.{child_name}" if prefix else child_name
# recurse into the child
self._sync_fsdp_params_to_vllm(child_module, prefix=child_prefix, visited=visited)
if isinstance(module, FSDP):
with FSDP.summon_full_params(module, recurse=False, writeback=False):
for param_name, param in module.named_parameters():
full_name = f"{prefix}.{param_name}" if prefix else param_name
for extra in ("_fsdp_wrapped_module.", "_checkpoint_wrapped_module."):
full_name = full_name.replace(extra, "")
if full_name in visited:
continue # skip FSDP subtrees already traversed
visited.add(full_name)
if self.vllm_mode == "server" and self.accelerator.is_main_process:
self.vllm_client.update_named_param(full_name, param.data)
elif self.vllm_mode == "colocate":
llm_model = (
self.vllm_engine.llm_engine.model_executor.driver_worker.model_runner.model
)
llm_model.load_weights([(full_name, param.data)])
def _move_model_to_vllm(self):
"""Synchronize student model weights to vLLM engine."""
# For DeepSpeed ZeRO-3 and FSDP, we need to gather all parameters before operations
deepspeed_plugin = self.accelerator.state.deepspeed_plugin
zero_stage_3 = deepspeed_plugin is not None and deepspeed_plugin.zero_stage == 3
if zero_stage_3:
import deepspeed
gather_if_zero3 = deepspeed.zero.GatheredParameters
else:
gather_if_zero3 = nullcontext
if self.vllm_mode == "colocate" and self.vllm_enable_sleep_mode:
empty_cache()
self.vllm_engine.wake_up(tags=["weights"])
if is_peft_model(self.model):
# With PEFT and FSDP/DeepSpeed ZeRO Stage 3, we must gather the full model at once before merging, as
# merging adapters in a sharded manner is not supported.
with gather_if_zero3(list(self.model.parameters())):
self.model.merge_adapter()
# Update vLLM weights while parameters are gathered
if self.is_fsdp_enabled: # note if using FSDP, gather_if_zero3 is nullcontext
# Update vLLM weights while parameters are gathered
# For PEFT with FSDP we need to use the memory efficient post-order traversal
self._sync_fsdp_params_to_vllm(self.model)
else:
# DeepSpeed ZeRO-3 with PEFT
for name, param in self.model.named_parameters():
# When using PEFT, we need to recover the original parameter name and discard some parameters
encrypted chain-of-thought — recorded, not exposed
args
{
"input": "const r = await tools.exec_command({\n cmd: \"sed -n '900,1115p' opsd_trainer.py\",\n workdir: \"/workspace/repo\",\n yield_time_ms: 10000,\n max_output_tokens: 14000\n});\ntext(r.output);"
}exec result
Script completed
Wall time 0.1 seconds
Output:
elapsed_time = time.time() - start_time
num_prompts = generated_tokens.shape[0]
total_completion_tokens = generated_tokens.shape[1] - inputs["student_prompts"].shape[1]
num_tokens = total_completion_tokens * num_prompts
avg_completion_length = total_completion_tokens
tokens_per_sec = num_tokens / elapsed_time if elapsed_time > 0 else 0
print(
f"generation done - elapsed time: {elapsed_time:.2f}s, prompts: {num_prompts}, total tokens: {num_tokens}, avg length: {avg_completion_length}, speed: {tokens_per_sec:.1f} tok/s"
)
new_attention_mask = torch.ones_like(generated_tokens)
new_labels = generated_tokens.clone()
if pad_token_id is not None:
new_labels[new_labels == pad_token_id] = -100
new_attention_mask[generated_tokens == pad_token_id] = 0
return generated_tokens, new_attention_mask, new_labels
@profiling_decorator
def _generate_on_policy_outputs_vllm(self, inputs, generation_config, pad_token_id=None):
"""Generate on-policy outputs from student prompts using vLLM."""
import time
device = self.accelerator.device
prompts_text_for_vllm = self.processing_class.batch_decode(
inputs["student_prompts"],
skip_special_tokens=False,
)
# Remove padding token text if it appears, as vLLM expects clean prompts
if self.processing_class.pad_token:
prompts_text_for_vllm = [
p.replace(self.processing_class.pad_token, "") for p in prompts_text_for_vllm
]
# Also decode prompts WITH special tokens for logging
prompts_text_with_special = self.processing_class.batch_decode(
inputs["student_prompts"],
skip_special_tokens=False,
)
# system_prompt = "Please reason step by step, and put your final answer within \\boxed{}."
# target_system_prompt = "You are Qwen, created by Alibaba Cloud. You are a helpful assistant."
# prompts_text = [p.replace(target_system_prompt, system_prompt) for p in prompts_text]
# Add system prompt to prompts
max_completion_length = generation_config.max_new_tokens
temperature = generation_config.temperature
# vLLM uses top_k=-1 for no top_k, transformers uses 0 or None.
top_k = generation_config.top_k if generation_config.top_k and generation_config.top_k > 0 else -1
# top_p, repetition_penalty, min_p, presence_penalty are not directly in generation_config, get from trainer args
top_p = self.args.top_p if hasattr(self.args, "top_p") else 1.0
repetition_penalty = self.args.repetition_penalty if hasattr(self.args, "repetition_penalty") else 1.0
min_p = self.args.min_p if hasattr(self.args, "min_p") else 0.0
presence_penalty = self.args.presence_penalty if hasattr(self.args, "presence_penalty") else 0.0
# Start timing for vLLM generation
start_time = time.time()
if self.vllm_mode == "server":
all_prompts_text = gather_object(prompts_text_for_vllm)
if self.accelerator.is_main_process:
completion_ids = self.vllm_client.generate(
prompts=all_prompts_text,
n=1, # In GKD, we generate 1 completion per prompt from student
repetition_penalty=repetition_penalty,
temperature=temperature,
top_p=top_p,
top_k=top_k,
min_p=min_p,
max_tokens=max_completion_length,
presence_penalty=presence_penalty,
guided_decoding_regex=self.vllm_guided_decoding_regex,
)
else:
completion_ids = [None] * len(all_prompts_text)
completion_ids = broadcast_object_list(completion_ids, from_process=0)
process_slice = slice(
self.accelerator.process_index * len(prompts_text_for_vllm),
(self.accelerator.process_index + 1) * len(prompts_text_for_vllm),
)
completion_ids = completion_ids[process_slice]
elif self.vllm_mode == "colocate":
if self.vllm_guided_decoding_regex:
guided_decoding = GuidedDecodingParams(
backend="outlines", regex=self.vllm_guided_decoding_regex
)
else:
guided_decoding = None
sampling_params = SamplingParams(
n=1,
repetition_penalty=repetition_penalty,
temperature=temperature,
top_p=top_p,
top_k=top_k,
min_p=min_p,
max_tokens=max_completion_length,
presence_penalty=presence_penalty,
guided_decoding=guided_decoding,
)
if hasattr(self, "vllm_tp_group") and self.vllm_tensor_parallel_size > 1:
# Gather prompts from all ranks in the TP group and flatten.
# Each rank starts with its own prompts; after gathering, all ranks see the full group set.
orig_size = len(prompts_text_for_vllm)
gathered_prompts = [None for _ in range(self.vllm_tensor_parallel_size)]
torch.distributed.all_gather_object(
gathered_prompts, prompts_text_for_vllm, group=self.vllm_tp_group
)
all_prompts_text = [p for sublist in gathered_prompts for p in sublist]
else:
all_prompts_text = prompts_text_for_vllm
all_outputs = self.vllm_engine.generate(
all_prompts_text, sampling_params=sampling_params, use_tqdm=False
)
completion_ids = [output.token_ids for outputs in all_outputs for output in outputs.outputs]
if hasattr(self, "vllm_tp_group") and self.vllm_tensor_parallel_size > 1:
# Slice completions for this rank within its TP group.
# Each rank generates all outputs — we keep only our share.
local_rank_in_group = torch.distributed.get_rank(group=self.vllm_tp_group)
tp_slice = slice(local_rank_in_group * orig_size, (local_rank_in_group + 1) * orig_size)
completion_ids = completion_ids[tp_slice]
if self.vllm_enable_sleep_mode:
self.vllm_engine.sleep(level=2)
else:
raise ValueError(f"Unknown vllm_mode: {self.vllm_mode}")
# Calculate and print vLLM generation statistics
elapsed_time = time.time() - start_time
total_completion_tokens = sum(len(ids) for ids in completion_ids)
num_prompts = len(completion_ids)
avg_completion_length = total_completion_tokens / num_prompts if num_prompts > 0 else 0
tokens_per_sec = total_completion_tokens / elapsed_time if elapsed_time > 0 else 0
print(
f"vLLM generation done - elapsed time: {elapsed_time:.2f}s, prompts: {num_prompts}, total tokens: {total_completion_tokens}, avg length: {avg_completion_length:.1f}, speed: {tokens_per_sec:.1f} tok/s"
)
# We need to combine prompt and completion for new_input_ids
# Tokenize prompts again to get prompt_ids on the correct device and format
# Use prompts_text_for_vllm (without special tokens) for tokenization since vLLM expects clean text
# Ensure add_special_tokens=False as vLLM typically handles prompts as raw text
# Calculate max_length for prompts, ensuring it's positive
prompt_max_length = (
max(1, self.args.max_length - max_completion_length) if self.args.max_length else None
)
prompt_tokenized = self.processing_class(
prompts_text_for_vllm,
return_tensors="pt",
padding="longest",
truncation=True if prompt_max_length else False,
max_length=prompt_max_length,
add_special_tokens=False,
).to(device)
prompt_ids = prompt_tokenized.input_ids
completion_ids_tensors = [torch.tensor(ids, device=device) for ids in completion_ids]
# Manually pad/truncate completions to max_completion_length length before using pad function
padded_completion_ids_list = []
for completion_tensor in completion_ids_tensors:
if len(completion_tensor) > max_completion_length:
# Truncate if longer than max_completion_length
padded_completion_ids_list.append(completion_tensor[:max_completion_length])
elif len(completion_tensor) < max_completion_length:
# Pad if shorter than max_completion_length
padding_needed = max_completion_length - len(completion_tensor)
padded_tensor = torch.cat(
[
completion_tensor,
torch.full(
(padding_needed,), pad_token_id, device=device, dtype=completion_tensor.dtype
),
]
)
padded_completion_ids_list.append(padded_tensor)
else:
# Already the right length
padded_completion_ids_list.append(completion_tensor)
# Now all tensors are the same length, so we can stack them
padded_completion_ids = torch.stack(padded_completion_ids_list)
# Ensure prompt_ids and padded_completion_ids are 2D
if prompt_ids.ndim == 1:
prompt_ids = prompt_ids.unsqueeze(0)
if padded_completion_ids.ndim == 1:
padded_completion_ids = padded_completion_ids.unsqueeze(0)
new_input_ids = torch.cat([prompt_ids, padded_completion_ids], dim=1)
new_attention_mask = torch.ones_like(new_input_ids, device=device)
new_labels = new_input_ids.clone()
if pad_token_id is not None:
new_labels[new_labels == pad_token_id] = -100
new_attention_mask[new_input_ids == pad_token_id] = 0
# Extract completion texts from the generated completion IDs
completion_texts = []
for comp_ids in completion_ids:
completion_text = self.processing_class.decode(comp_ids, skip_special_tokens=False)
completion_texts.append(completion_text)
return new_input_ids, new_attention_mask, new_labels, prompts_text_with_special, completion_texts
def _generate_teacher_reasoning_vllm(
self, teacher_reasoning_prompts, teacher_reasoning_attention_mask=None
):
"""Generate teacher's reasoning using vLLM."""
import time
device = self.accelerator.device
encrypted chain-of-thought — recorded, not exposed
args
{
"input": "const r = await tools.update_plan({explanation:\"Implemented and validated the selected OPSD improvement. Three independent official-style 100-step train + full AIME24 evaluations average 55.00%. Final syntax, diff, recipe, and frozen-cache integrity checks pass.\",\"plan\":[{\"step\":\"Inspect the paper, baseline implementation, logs, and dataset characteristics\",\"status\":\"completed\"},{\"step\":\"Implement the proposed OPSD method improvement\",\"status\":\"completed\"},{\"step\":\"Run capped training and official-style AIME24 evaluations\",\"status\":\"completed\"},{\"step\":\"Compare ablations and select the strongest robust candidate\",\"status\":\"completed\"},{\"step\":\"Restore frozen inputs and perform final reproducibility checks\",\"status\":\"completed\"}]});\ntext(r);"
}exec result
Script completed
Wall time 0.0 seconds
Output:
{}Implemented and validated the improved OPSD method.
Key changes:
- AIME-aligned olympiad/AoPS integer-answer curriculum.
- Early-token-weighted OPSD loss.
- Correctness-gated on-policy rejection-SFT auxiliary loss.
- Linear LR schedule with
lmbda=0.01.
Official-style full evaluations: 56.39, 53.89, 54.72, averaging 55.00%, versus the 54.22% baseline.
Changed files:
Syntax, diff, and frozen-input integrity checks pass. The pre-existing README modification was left untouched.