diff --git a/batch_runner.py b/batch_runner.py index 4b5006ef77..9ceac8f5c8 100644 --- a/batch_runner.py +++ b/batch_runner.py @@ -46,6 +46,13 @@ ALL_POSSIBLE_TOOLS = set(TOOL_TO_TOOLSET_MAP.keys()) DEFAULT_TOOL_STATS = {'count': 0, 'success': 0, 'failure': 0} _REASONING_KEYS = ("total_assistant_turns", "turns_with_reasoning", "turns_without_reasoning") +# BatchRunner.__init__ parameters stored as same-named attributes. +_RUNNER_FIELDS = ( + "batch_size", "run_name", "distribution", "max_iterations", "base_url", "api_key", "model", + "num_workers", "verbose", "ephemeral_system_prompt", "log_prefix_chars", "providers_allowed", + "providers_ignored", "providers_order", "provider_sort", "openrouter_min_coding_score", + "max_tokens", "reasoning_config", "prefill_messages", "max_samples", +) # BatchRunner attributes forwarded verbatim to every AIAgent in the worker config. _AGENT_PASSTHROUGH = ( "base_url", "api_key", "ephemeral_system_prompt", "providers_allowed", "providers_ignored", @@ -130,24 +137,18 @@ def _extract_tool_stats(messages: List[Dict[str, Any]]) -> Dict[str, Dict[str, i return tool_stats +def _turn_has_reasoning(msg: Dict[str, Any]) -> bool: + """```` in content, or a non-empty native ``reasoning`` field.""" + if "" in (msg.get("content", "") or ""): + return True + return bool(msg.get("reasoning", "").strip()) if msg.get("reasoning") else False + + def _extract_reasoning_stats(messages: List[Dict[str, Any]]) -> Dict[str, int]: - """Count assistant turns with reasoning (```` in content or a - non-empty native ``reasoning`` field) vs without.""" - total = 0 - with_reasoning = 0 - - for msg in messages: - if msg.get("role") != "assistant": - continue - total += 1 - - content = msg.get("content", "") or "" - has_scratchpad = "" in content - has_native_reasoning = bool(msg.get("reasoning", "").strip()) if msg.get("reasoning") else False - - if has_scratchpad or has_native_reasoning: - with_reasoning += 1 - + """Count assistant turns with reasoning vs without.""" + assistant_turns = [msg for msg in messages if msg.get("role") == "assistant"] + total = len(assistant_turns) + with_reasoning = sum(1 for msg in assistant_turns if _turn_has_reasoning(msg)) return { "total_assistant_turns": total, "turns_with_reasoning": with_reasoning, @@ -405,6 +406,12 @@ def _entry_prompt_text(entry: Dict) -> str: return "" +def _banner(title: str) -> None: + print("\n" + "=" * 70) + print(title) + print("=" * 70) + + def _chunk(entries: List[Tuple[int, Dict[str, Any]]], size: int) -> List[List[Tuple[int, Dict[str, Any]]]]: """Split ``(index, entry)`` tuples into batches of *size*, preserving original indices.""" return [entries[i:i + size] for i in range(0, len(entries), size)] @@ -443,27 +450,10 @@ class BatchRunner: ``prefill_messages`` are prepended as few-shot context; Anthropic Sonnet/Opus 4.6+ reject a trailing assistant-role prefill (400) — use user-role priming for those. """ + params = dict(locals()) self.dataset_file = Path(dataset_file) - self.batch_size = batch_size - self.run_name = run_name - self.distribution = distribution - self.max_iterations = max_iterations - self.base_url = base_url - self.api_key = api_key - self.model = model - self.num_workers = num_workers - self.verbose = verbose - self.ephemeral_system_prompt = ephemeral_system_prompt - self.log_prefix_chars = log_prefix_chars - self.providers_allowed = providers_allowed - self.providers_ignored = providers_ignored - self.providers_order = providers_order - self.provider_sort = provider_sort - self.openrouter_min_coding_score = openrouter_min_coding_score - self.max_tokens = max_tokens - self.reasoning_config = reasoning_config - self.prefill_messages = prefill_messages - self.max_samples = max_samples + for name in _RUNNER_FIELDS: + setattr(self, name, params[name]) if not validate_distribution(distribution): raise ValueError(f"Unknown distribution: {distribution}. Available: {list(list_distributions().keys())}") @@ -627,9 +617,7 @@ class BatchRunner: self.batches = _chunk(filtered_entries, self.batch_size) - print("\n" + "=" * 70) - print("📊 RESUME SUMMARY") - print("=" * 70) + _banner("📊 RESUME SUMMARY") print(f" Original dataset size: {len(self.dataset):,} prompts") print(f" Already completed: {len(skipped_indices):,} prompts") print(" ─────────────────────────────────────────") @@ -781,9 +769,7 @@ class BatchRunner: return kept, batch_files_found def _print_summary(self, results, total_tool_stats, total_reasoning_stats, kept, batch_files_found, start_time) -> None: - print("\n" + "=" * 70) - print("📊 BATCH PROCESSING COMPLETE") - print("=" * 70) + _banner("📊 BATCH PROCESSING COMPLETE") print(f"✅ Prompts processed this run: {sum(r.get('processed', 0) for r in results)}") print(f"✅ Total trajectories in merged file: {kept}") print(f"✅ Total batch files merged: {batch_files_found}") @@ -827,9 +813,7 @@ class BatchRunner: def run(self, resume: bool = False): """Run the batch pipeline; with *resume*, skip prompts already present in batch files.""" - print("\n" + "=" * 70) - print("🚀 Starting Batch Processing") - print("=" * 70) + _banner("🚀 Starting Batch Processing") if resume and not self._apply_resume(): return