Skip to content

Commit c2e1438

Browse files
feat: implement the beam search algorithm
1 parent 06af4f0 commit c2e1438

4 files changed

Lines changed: 389 additions & 62 deletions

File tree

Lines changed: 192 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,192 @@
1+
import asyncio
2+
import logging
3+
import random
4+
from typing import List, Tuple
5+
6+
from auto_search.graph import Node
7+
from auto_search.orchestrator import SearchOrchestrator
8+
from auto_search.strategies import TPU_PALLAS_OPTIMIZATION_STRATEGIES
9+
from auto_search.worker import ADKSessionWorker
10+
11+
logger = logging.getLogger(__name__)
12+
13+
14+
class BeamSearchOrchestrator(SearchOrchestrator):
15+
def __init__(
16+
self,
17+
beam_size: int = 2,
18+
branches_per_node: int = 2,
19+
max_depth: int = 2,
20+
keep_factor: float = 1,
21+
strategies: List[str] = TPU_PALLAS_OPTIMIZATION_STRATEGIES,
22+
agent_config: dict = None,
23+
**kwargs,
24+
):
25+
self._validate_args(beam_size, branches_per_node, max_depth, keep_factor)
26+
self.beam_size = beam_size
27+
self.branches_per_node = branches_per_node
28+
self.strategies = strategies
29+
self.max_depth = max_depth
30+
self.keep_factor = keep_factor
31+
self.agent_config = agent_config or {"max_iterations": 1}
32+
self.worker = ADKSessionWorker()
33+
34+
self.current_depth = 0
35+
self.beam: List[Node] = []
36+
37+
super().__init__(**kwargs)
38+
39+
if not self.graph.metadata:
40+
self.beam = [self.graph.get_node(self.graph.root_id)]
41+
self.update_metadata("current_depth", self.current_depth)
42+
self.update_metadata(
43+
"beam_node_ids", [node.node_id for node in self.beam]
44+
)
45+
46+
def _validate_args(
47+
self,
48+
beam_size: int,
49+
branches_per_node: int,
50+
max_depth: int,
51+
keep_factor: float,
52+
) -> None:
53+
if beam_size < 1:
54+
raise ValueError(f"beam_size must be at least 1, got {beam_size}.")
55+
if branches_per_node < 1:
56+
raise ValueError(
57+
f"branches_per_node must be at least 1, got {branches_per_node}."
58+
)
59+
if max_depth < 1:
60+
raise ValueError(f"max_depth must be at least 1, got {max_depth}.")
61+
if keep_factor <= 0:
62+
raise ValueError(f"keep_factor must be positive, got {keep_factor}.")
63+
64+
def _resume(self) -> None:
65+
self.current_depth = self.graph.metadata.get("current_depth", 0)
66+
beam_ids = self.graph.metadata.get("beam_node_ids", [])
67+
self.beam = [
68+
node
69+
for node in (self.graph.get_node(node_id) for node_id in beam_ids)
70+
if node is not None
71+
]
72+
logger.info(
73+
f"Resumed Beam Search at depth {self.current_depth} with beam: {beam_ids}"
74+
)
75+
76+
def _select_nodes_to_expand(self) -> List[Node]:
77+
return self.beam
78+
79+
def _generate_expansion_tasks(
80+
self, nodes: List[Node]
81+
) -> List[Tuple[Node, str]]:
82+
tasks = []
83+
for node in nodes:
84+
selected_strategies = random.sample(
85+
self.strategies,
86+
min(self.branches_per_node, len(self.strategies)),
87+
)
88+
for strategy in selected_strategies:
89+
tasks.append((node, strategy))
90+
return tasks
91+
92+
def _update_search_state(self, new_nodes: List[Node]) -> None:
93+
candidates = []
94+
regressed_candidates = []
95+
for node in new_nodes:
96+
if not node.is_valid_candidate:
97+
logger.warning(
98+
f"Node {node.node_id} failed Validity Check. "
99+
"Adding to regressed candidates"
100+
)
101+
regressed_candidates.append(node)
102+
continue
103+
104+
parent = self.graph.get_node(node.parent_id)
105+
parent_latency = parent.evaluation.latency_ms
106+
parent_latency = (
107+
parent_latency if parent_latency is not None else float("inf")
108+
)
109+
110+
current_latency = node.evaluation.latency_ms
111+
current_latency = (
112+
current_latency if current_latency is not None else float("inf")
113+
)
114+
115+
if current_latency < parent_latency * self.keep_factor:
116+
candidates.append(node)
117+
else:
118+
logger.info(
119+
f"Node {node.node_id} failed Parent Regression Gate. "
120+
"Adding to regressed candidates"
121+
)
122+
regressed_candidates.append(node)
123+
124+
candidates.sort(key=lambda n: n.evaluation.latency_ms)
125+
126+
if len(candidates) < self.beam_size and regressed_candidates:
127+
shortage = self.beam_size - len(candidates)
128+
logger.warning(
129+
f"Only {len(candidates)} candidates passed the Parent Regression Gate. "
130+
f"Padding with the best {shortage} regressed candidates to keep search alive."
131+
)
132+
regressed_candidates.sort(
133+
key=lambda n: (
134+
n.evaluation.latency_ms
135+
if n.evaluation.latency_ms is not None
136+
else float("inf")
137+
)
138+
)
139+
candidates.extend(regressed_candidates)
140+
141+
self.beam = candidates[: self.beam_size]
142+
self.update_metadata("beam_node_ids", [n.node_id for n in self.beam])
143+
144+
def _post_step_hook(self) -> None:
145+
self.current_depth += 1
146+
self.update_metadata("current_depth", self.current_depth)
147+
148+
def _should_terminate(self) -> bool:
149+
return self.current_depth >= self.max_depth or not self.beam
150+
151+
async def _execute_expansions(
152+
self, tasks: List[Tuple[Node, str]]
153+
) -> List[Node]:
154+
async def run_task(task_idx: int, parent_node: Node, strategy: str) -> Node:
155+
node_id, base_dir = self.get_next_session_node()
156+
for attempt in range(1, self.max_worker_retries + 1):
157+
logger.info(
158+
f"Task {task_idx}: Expanding {parent_node.node_id} using \n"
159+
f" strategy '{strategy}' -> {node_id}. \n"
160+
f"Attempt {attempt}/{self.max_worker_retries}."
161+
)
162+
async with self._semaphore:
163+
session_dir = f"{base_dir}_attempt_{attempt}"
164+
node = await self.worker.expand_node(
165+
node_id,
166+
parent_node,
167+
session_dir=session_dir,
168+
reference_code=self.reference_code,
169+
strategy=strategy,
170+
agent_config=self.agent_config,
171+
)
172+
if (
173+
node.execution_status == "SUCCESS"
174+
or attempt == self.max_worker_retries
175+
):
176+
self.graph.add_node(node)
177+
logger.info(
178+
f"Task {task_idx}: Finished {node_id} with status {node.execution_status} "
179+
f"(Latency: {node.evaluation.latency_ms} ms)"
180+
)
181+
return node
182+
183+
logger.warning(
184+
f"Task {task_idx} (strategy: {strategy}) failed attempt"
185+
f" {attempt}/{self.max_worker_retries}: {node.execution_error}. Retrying..."
186+
)
187+
await asyncio.sleep(2**attempt)
188+
189+
futures = [
190+
run_task(i, parent, strat) for i, (parent, strat) in enumerate(tasks)
191+
]
192+
return await asyncio.gather(*futures)

MaxKernel/auto_search/run_batch_search.py

Lines changed: 51 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -40,9 +40,7 @@ async def run_batch_search(
4040
):
4141
"""Coordinates concurrent problem execution across the dataset."""
4242
if not os.path.isdir(data_dir):
43-
logger.error(
44-
f"Dataset directory not found or not a directory: {data_dir}"
45-
)
43+
logger.error(f"Dataset directory not found or not a directory: {data_dir}")
4644
return
4745

4846
data_dir_valid = [
@@ -135,36 +133,65 @@ def parse_args() -> argparse.Namespace:
135133
default=None,
136134
help="File to save logs to",
137135
)
138-
# Parallel Search Arguments
139-
parallel_group = parser.add_argument_group(
140-
"Parallel Search Arguments",
141-
"Parameters specific to the 'parallel' search algorithm.",
142-
)
143-
parallel_group.add_argument(
144-
"--num_parallel_runs",
145-
type=int,
146-
default=2,
147-
help="Number of parallel runs",
148-
)
149-
parallel_group.add_argument(
150-
"--max_retries",
136+
orch_group.add_argument(
137+
"--max_worker_retries",
151138
type=int,
152139
default=1,
153140
help="Max worker retries per expansion task",
154141
)
155-
parallel_group.add_argument(
142+
orch_group.add_argument(
156143
"--strategies",
157144
nargs="+",
158145
type=str,
159146
default=None,
160147
help="List of strategy strings to explore",
161148
)
162-
parallel_group.add_argument(
149+
orch_group.add_argument(
163150
"--agent_config",
164151
type=str,
165152
default=None,
166153
help="JSON string of agent config parameters (e.g. '{\"max_iterations\": 5}')",
167154
)
155+
# Parallel Search Arguments
156+
parallel_group = parser.add_argument_group(
157+
"Parallel Search Arguments",
158+
"Parameters specific to the 'parallel' search algorithm.",
159+
)
160+
parallel_group.add_argument(
161+
"--num_parallel_runs",
162+
type=int,
163+
default=2,
164+
help="Number of parallel runs",
165+
)
166+
# Beam Search Arguments
167+
beam_group = parser.add_argument_group(
168+
"Beam Search Arguments",
169+
"Parameters specific to the 'beam' search algorithm.",
170+
)
171+
beam_group.add_argument(
172+
"--beam_size",
173+
type=int,
174+
default=2,
175+
help="Size of the beam (number of candidates to keep per depth)",
176+
)
177+
beam_group.add_argument(
178+
"--branches_per_node",
179+
type=int,
180+
default=2,
181+
help="Number of branches/strategies to explore per node in the beam",
182+
)
183+
beam_group.add_argument(
184+
"--max_depth",
185+
type=int,
186+
default=2,
187+
help="Maximum depth of the beam search",
188+
)
189+
beam_group.add_argument(
190+
"--keep_factor",
191+
type=float,
192+
default=1.0,
193+
help="Factor of parent latency to keep candidates (e.g. 1.0 means must not be worse than parent)",
194+
)
168195
return parser.parse_args()
169196

170197

@@ -182,10 +209,14 @@ def main():
182209

183210
kwargs = {
184211
"max_concurrency": args.max_concurrency,
185-
"num_parallel_runs": args.num_parallel_runs,
186-
"max_worker_retries": args.max_retries,
212+
"max_worker_retries": args.max_worker_retries,
187213
"strategies": args.strategies,
188214
"agent_config": parsed_agent_config,
215+
"num_parallel_runs": args.num_parallel_runs,
216+
"beam_size": args.beam_size,
217+
"branches_per_node": args.branches_per_node,
218+
"max_depth": args.max_depth,
219+
"keep_factor": args.keep_factor,
189220
}
190221

191222
asyncio.run(

0 commit comments

Comments
 (0)