Skip to content

Commit a67bf94

Browse files
feat: implement the parallel search algorithm
1 parent 15bf797 commit a67bf94

1 file changed

Lines changed: 145 additions & 0 deletions

File tree

Lines changed: 145 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,145 @@
1+
import asyncio
2+
import logging
3+
from typing import List, Optional, Tuple
4+
5+
from auto_search.graph import Node
6+
from auto_search.orchestrator import SearchOrchestrator
7+
from auto_search.worker import ADKSessionWorker
8+
9+
logger = logging.getLogger(__name__)
10+
11+
12+
class SimpleParallelSearchOrchestrator(SearchOrchestrator):
13+
"""Orchestrates single-iteration parallel kernel explorations from root."""
14+
15+
def __init__(
16+
self,
17+
num_parallel_runs: int = 2,
18+
strategies: Optional[List[str]] = None,
19+
max_worker_retries: int = 1,
20+
agent_config: Optional[dict] = None,
21+
**kwargs,
22+
):
23+
if num_parallel_runs <= 0:
24+
raise ValueError(
25+
f"num_parallel_runs must be a positive integer, got {num_parallel_runs}."
26+
)
27+
if max_worker_retries < 1:
28+
raise ValueError(
29+
f"max_worker_retries must be at least 1, got {max_worker_retries}."
30+
)
31+
32+
self.num_parallel_runs = num_parallel_runs
33+
if strategies is not None:
34+
if len(strategies) > num_parallel_runs:
35+
raise ValueError(
36+
f"Number of user specified strategies ({len(strategies)}) cannot"
37+
f" exceed num_parallel_runs ({num_parallel_runs})."
38+
)
39+
self.strategies = [s or "" for s in strategies] + [""] * (
40+
num_parallel_runs - len(strategies)
41+
)
42+
else:
43+
self.strategies = [""] * num_parallel_runs
44+
45+
self.remaining_strategies = list(self.strategies)
46+
self.max_worker_retries = max_worker_retries
47+
self.agent_config = agent_config
48+
self.worker = ADKSessionWorker()
49+
50+
super().__init__(**kwargs)
51+
52+
def _resume(self) -> None:
53+
if not self.graph.root_id:
54+
logger.warning("Cannot resume: root_id is not set in graph.")
55+
return
56+
57+
existing_children = [
58+
node
59+
for node in self.graph.nodes.values()
60+
if node.parent_id == self.graph.root_id
61+
]
62+
63+
for node in existing_children:
64+
strategy = node.strategy_applied or ""
65+
if strategy in self.remaining_strategies:
66+
self.remaining_strategies.remove(strategy)
67+
logger.info(
68+
f"Resumed Simple Parallel Search. Found {len(existing_children)} existing"
69+
f" evaluations. {len(self.remaining_strategies)} runs remaining."
70+
)
71+
72+
def _select_nodes_to_expand(self) -> List[Node]:
73+
if self._should_terminate():
74+
return []
75+
if not self.graph.root_id:
76+
logger.error(
77+
"Cannot select nodes to expand: root_id is not set in graph."
78+
)
79+
return []
80+
root_node = self.graph.get_node(self.graph.root_id)
81+
return [root_node] if root_node else []
82+
83+
def _generate_expansion_tasks(
84+
self, nodes: List[Node]
85+
) -> List[Tuple[Node, str]]:
86+
if not self.remaining_strategies:
87+
logger.info(
88+
f"All {self.num_parallel_runs} parallel runs have already completed."
89+
)
90+
return []
91+
92+
root_node = self.graph.get_node(self.graph.root_id)
93+
if not root_node:
94+
return []
95+
96+
logger.info(
97+
f"Launching {len(self.remaining_strategies)} parallel runs for root node."
98+
)
99+
return [(root_node, strategy) for strategy in self.remaining_strategies]
100+
101+
def _update_search_state(self, new_nodes: List[Node]) -> None:
102+
for node in new_nodes:
103+
strategy = node.strategy_applied or ""
104+
if strategy in self.remaining_strategies:
105+
self.remaining_strategies.remove(strategy)
106+
107+
def _should_terminate(self) -> bool:
108+
return len(self.remaining_strategies) == 0
109+
110+
async def _execute_expansions(
111+
self, tasks: List[Tuple[Node, str]]
112+
) -> List[Node]:
113+
async def run_task(task_idx: int, parent_node: Node, strategy: str) -> Node:
114+
node_id, base_dir = self.get_next_session_node()
115+
for attempt in range(1, self.max_worker_retries + 1):
116+
async with self._semaphore:
117+
session_dir = f"{base_dir}_attempt_{attempt}"
118+
node = await self.worker.expand_node(
119+
node_id,
120+
parent_node,
121+
strategy=strategy,
122+
session_dir=session_dir,
123+
reference_code=self.reference_code,
124+
agent_config=self.agent_config,
125+
)
126+
if (
127+
node.execution_status == "SUCCESS"
128+
or attempt == self.max_worker_retries
129+
):
130+
self.graph.add_node(node)
131+
return node
132+
133+
strat_info = f" (strategy: {strategy})" if strategy else ""
134+
logger.warning(
135+
f"Task {task_idx}{strat_info} failed attempt"
136+
f" {attempt}/{self.max_worker_retries}: {node.execution_error}."
137+
" Retrying..."
138+
)
139+
await asyncio.sleep(2**attempt)
140+
141+
futures = [
142+
run_task(i, parent, strategy)
143+
for i, (parent, strategy) in enumerate(tasks)
144+
]
145+
return await asyncio.gather(*futures)

0 commit comments

Comments
 (0)