diff --git a/batch_runner.py b/batch_runner.py index af6dd75c1e919..5eccb1a53aa26 100644 --- a/batch_runner.py +++ b/batch_runner.py @@ -457,6 +457,17 @@ def _process_batch_worker(args: Tuple) -> Dict[str, Any]: print(f" đŸšĢ Prompt {prompt_index} discarded (no reasoning in any turn)") discarded_no_reasoning += 1 completed_in_batch.append(prompt_index) + # Write a tombstone row so the content-based resume scan (which + # only reads batch_*.jsonl) can see this prompt was already + # processed and discarded, not just left unprocessed. + with open(batch_output_file, 'a', encoding='utf-8') as f: + f.write(json.dumps({ + "prompt_index": prompt_index, + "conversations": result["trajectory"], + "discarded": "no_reasoning", + }, ensure_ascii=False) + "\n") + f.flush() + os.fsync(f.fileno()) continue # Get and normalize tool stats for consistent schema across all entries @@ -995,8 +1006,11 @@ def run(self, resume: bool = False): # Aggregate all batch statistics and update checkpoint total_reasoning_stats = {"total_assistant_turns": 0, "turns_with_reasoning": 0, "turns_without_reasoning": 0} + total_discarded_no_reasoning = 0 for batch_result in results: + total_discarded_no_reasoning += batch_result.get("discarded_no_reasoning", 0) + # Aggregate tool stats for tool_name, stats in batch_result.get("tool_stats", {}).items(): if tool_name not in total_tool_stats: @@ -1043,6 +1057,7 @@ def run(self, resume: bool = False): total_entries = 0 filtered_entries = 0 + discarded_tombstones = 0 batch_files_found = 0 # Find ALL batch files in the output directory (handles resume merging old + new) @@ -1058,8 +1073,16 @@ def run(self, resume: bool = False): total_entries += 1 try: data = json.loads(line) + + # Discard tombstones exist only so resume can see + # these prompts as done; they carry no full + # trajectory and must not enter the training file. + if data.get("discarded"): + discarded_tombstones += 1 + continue + tool_stats = data.get('tool_stats', {}) - + # Check for invalid tool names (model hallucinations) invalid_tools = [k for k in tool_stats if k not in VALID_TOOLS] @@ -1076,7 +1099,9 @@ def run(self, resume: bool = False): if filtered_entries > 0: print(f"âš ī¸ Filtered {filtered_entries} corrupted entries out of {total_entries} total") - print(f"✅ Combined {batch_files_found} batch files into trajectories.jsonl ({total_entries - filtered_entries} entries)") + if discarded_tombstones > 0: + print(f"â„šī¸ Excluded {discarded_tombstones} discarded (no-reasoning) tombstone rows out of {total_entries} total") + print(f"✅ Combined {batch_files_found} batch files into trajectories.jsonl ({total_entries - filtered_entries - discarded_tombstones} entries)") # Save final statistics final_stats = { @@ -1090,6 +1115,7 @@ def run(self, resume: bool = False): "duration_seconds": round(time.time() - start_time, 2), "tool_statistics": total_tool_stats, "reasoning_statistics": total_reasoning_stats, + "discarded_no_reasoning": total_discarded_no_reasoning, } with open(self.stats_file, 'w', encoding='utf-8') as f: diff --git a/tests/test_batch_runner_checkpoint.py b/tests/test_batch_runner_checkpoint.py index d2fb370b38f41..0a9096b3a3565 100644 --- a/tests/test_batch_runner_checkpoint.py +++ b/tests/test_batch_runner_checkpoint.py @@ -153,7 +153,8 @@ def test_discarded_no_reasoning_prompts_are_marked_completed(self, tmp_path, mon batch_file = tmp_path / "batch_1.jsonl" prompt_result = { "success": True, - "trajectory": [{"role": "assistant", "content": "x"}], + "trajectory": [{"from": "human", "value": "hi"}, + {"role": "assistant", "content": "x"}], "reasoning_stats": {"has_any_reasoning": False}, "tool_stats": {}, "metadata": {}, @@ -174,7 +175,52 @@ def test_discarded_no_reasoning_prompts_are_marked_completed(self, tmp_path, mon assert result["discarded_no_reasoning"] == 1 assert result["completed_prompts"] == [0] - assert not batch_file.exists() or batch_file.read_text() == "" + + # A tombstone row must be written so the content-based resume scan + # can see this prompt was already processed and discarded. + assert batch_file.exists() + lines = [l for l in batch_file.read_text(encoding="utf-8").strip().split("\n") if l] + assert len(lines) == 1 + entry = json.loads(lines[0]) + assert entry["discarded"] == "no_reasoning" + + def test_resume_after_all_discarded_batch_reruns_zero_prompts(self, tmp_path, monkeypatch): + """Regression for the issue: a resumed run must not re-execute + prompts that were already processed and discarded for having no + reasoning — the content-based scan must see the discard tombstone. + """ + prompt_result = { + "success": True, + "trajectory": [{"from": "human", "value": "hi"}, + {"role": "assistant", "content": "x"}], + "reasoning_stats": {"has_any_reasoning": False}, + "tool_stats": {}, + "metadata": {}, + "completed": True, + "api_calls": 1, + "toolsets_used": [], + } + monkeypatch.setattr("batch_runner._process_single_prompt", lambda *args, **kwargs: prompt_result) + + # First run: prompt 0 gets processed and discarded, writing its + # tombstone row into batch_1.jsonl. + _process_batch_worker((1, [(0, {"prompt": "hi"})], tmp_path, set(), {"verbose": False})) + + # Simulate a fresh resume: scan batch files by content, exactly as + # BatchRunner.run() does. + r = BatchRunner.__new__(BatchRunner) + r.output_dir = tmp_path + completed_prompt_texts = r._scan_completed_prompts_by_content() + + assert "hi" in completed_prompt_texts, ( + "discarded prompt is invisible to the content-based resume scan" + ) + + r.dataset = [{"prompt": "hi"}] + filtered_entries, skipped_indices = r._filter_dataset_by_completed(completed_prompt_texts) + + assert filtered_entries == [], "discarded prompt was rescheduled on resume" + assert skipped_indices == [0] class TestFinalCheckpointNoDuplicates: