Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 28 additions & 2 deletions batch_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand All @@ -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]

Expand All @@ -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 = {
Expand All @@ -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:
Expand Down
50 changes: 48 additions & 2 deletions tests/test_batch_runner_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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": {},
Expand All @@ -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:
Expand Down
Loading