1use super::{
2 AgentSimulator, PerceptMap, PerceptOutcome, best_action_from_action_values,
3 choose_uniform_unvisited, ensure_action_slots, prune_key, random_rollout,
4};
5#[cfg(test)]
6use crate::aixi::common::ActionAlphabet;
7use crate::aixi::common::{Action, PerceptVal, Reward};
8use rayon::prelude::*;
9use std::collections::HashMap;
10use std::fmt;
11use std::num::NonZeroUsize;
12#[cfg(test)]
13use std::sync::{Arc, Mutex};
14
15const PARALLEL_PLANNER_SEED_SALT: u64 = 0x9E37_79B9_7F4A_7C15;
16
17#[derive(Clone, Copy, Debug, Eq, PartialEq)]
23pub enum ParallelUctPlannerInitError {
24 InvalidBuUctMMax,
26}
27
28impl fmt::Display for ParallelUctPlannerInitError {
29 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
30 match self {
31 Self::InvalidBuUctMMax => write!(f, "parallel_uct bu_uct_m_max must be in (0, 1)"),
32 }
33 }
34}
35
36impl std::error::Error for ParallelUctPlannerInitError {}
37
38#[derive(Clone, Copy, Debug, Eq, PartialEq)]
49pub enum ParallelUctSearchError {
50 PositiveSamplesRequirePositiveHorizon,
56}
57
58impl fmt::Display for ParallelUctSearchError {
59 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
60 match self {
61 Self::PositiveSamplesRequirePositiveHorizon => {
62 write!(
63 f,
64 "parallel_uct requires agent.horizon() >= 1 when samples > 0"
65 )
66 }
67 }
68 }
69}
70
71impl std::error::Error for ParallelUctSearchError {}
72
73pub struct ParallelUctPlanner {
79 state: ParallelPlannerState,
80 workers: NonZeroUsize,
81 bu_uct_m_max: Option<f64>,
82}
83
84impl ParallelUctPlanner {
85 pub fn new(
90 workers: NonZeroUsize,
91 bu_uct_m_max: Option<f64>,
92 ) -> Result<Self, ParallelUctPlannerInitError> {
93 if matches!(bu_uct_m_max, Some(m_max) if !(0.0 < m_max && m_max < 1.0)) {
94 return Err(ParallelUctPlannerInitError::InvalidBuUctMMax);
95 }
96 Ok(Self {
97 state: if bu_uct_m_max.is_some() {
98 ParallelPlannerState::Bu(ParallelRuntime::new())
99 } else {
100 ParallelPlannerState::Wu(ParallelRuntime::new())
101 },
102 workers,
103 bu_uct_m_max,
104 })
105 }
106
107 pub fn search(
153 &mut self,
154 agent: &mut dyn AgentSimulator,
155 prev_obs_stream: &[PerceptVal],
156 prev_rew: Reward,
157 prev_act: Action,
158 samples: usize,
159 ) -> Result<Action, ParallelUctSearchError> {
160 let horizon = agent.horizon();
161 if samples > 0 && horizon == 0 {
162 return Err(ParallelUctSearchError::PositiveSamplesRequirePositiveHorizon);
163 }
164 Ok(self.search_validated_with_horizon(
165 agent,
166 prev_obs_stream,
167 prev_rew,
168 prev_act,
169 samples,
170 horizon,
171 ))
172 }
173
174 pub(crate) fn search_validated(
175 &mut self,
176 agent: &mut dyn AgentSimulator,
177 prev_obs_stream: &[PerceptVal],
178 prev_rew: Reward,
179 prev_act: Action,
180 samples: usize,
181 ) -> Action {
182 let horizon = agent.horizon();
183 debug_assert!(
184 samples == 0 || horizon > 0,
185 "parallel_uct validated search requires agent.horizon() >= 1 when samples > 0"
186 );
187 self.search_validated_with_horizon(
188 agent,
189 prev_obs_stream,
190 prev_rew,
191 prev_act,
192 samples,
193 horizon,
194 )
195 }
196
197 fn search_validated_with_horizon(
198 &mut self,
199 agent: &mut dyn AgentSimulator,
200 prev_obs_stream: &[PerceptVal],
201 prev_rew: Reward,
202 prev_act: Action,
203 samples: usize,
204 horizon: usize,
205 ) -> Action {
206 let workers = self.workers.get();
207 match &mut self.state {
208 ParallelPlannerState::Wu(runtime) => search_runtime::<WuMode>(
209 runtime,
210 agent,
211 SearchRuntimeParams {
212 prev_obs_stream,
213 prev_rew,
214 prev_act,
215 samples,
216 horizon,
217 workers,
218 bu_uct_m_max: None,
219 },
220 ),
221 ParallelPlannerState::Bu(runtime) => search_runtime::<BuMode>(
222 runtime,
223 agent,
224 SearchRuntimeParams {
225 prev_obs_stream,
226 prev_rew,
227 prev_act,
228 samples,
229 horizon,
230 workers,
231 bu_uct_m_max: self.bu_uct_m_max,
232 },
233 ),
234 }
235 }
236}
237
238struct SearchRuntimeParams<'a> {
239 prev_obs_stream: &'a [PerceptVal],
240 prev_rew: Reward,
241 prev_act: Action,
242 samples: usize,
243 horizon: usize,
244 workers: usize,
245 bu_uct_m_max: Option<f64>,
246}
247
248fn search_runtime<M: ModeState>(
249 runtime: &mut ParallelRuntime<M>,
250 agent: &mut dyn AgentSimulator,
251 params: SearchRuntimeParams<'_>,
252) -> Action {
253 let SearchRuntimeParams {
254 prev_obs_stream,
255 prev_rew,
256 prev_act,
257 samples,
258 horizon,
259 workers,
260 bu_uct_m_max,
261 } = params;
262 prune_tree(runtime, agent, prev_obs_stream, prev_rew, prev_act);
263
264 debug_assert!(workers > 0);
265 let logical_workers = workers.min(samples.max(1));
266 let planner_seed = agent.gen_f64().to_bits();
267 let gamma = agent.discount_gamma().clamp(0.0, 1.0);
268
269 let mut dispatched = 0usize;
270 if samples > 0 {
271 let root_is_fresh = runtime.root.as_ref().is_some_and(|root| root.visits == 0);
272 if root_is_fresh {
273 let task_index = 0usize;
274 let mut local_agent =
275 agent.boxed_clone_with_seed(planner_task_seed(planner_seed, task_index));
276 local_agent.begin_discardable_simulation();
277 let bootstrap = bootstrap_root(runtime, local_agent.as_mut(), horizon, task_index);
278 complete_update_batch(runtime, std::slice::from_ref(&bootstrap), gamma);
279 dispatched = 1;
280 }
281 }
282
283 while dispatched < samples {
284 let batch_size = logical_workers.min(samples - dispatched);
285 let mut pending = Vec::with_capacity(batch_size);
286
287 for batch_index in 0..batch_size {
288 let task_index = dispatched + batch_index;
289 let mut local_agent =
290 agent.boxed_clone_with_seed(planner_task_seed(planner_seed, task_index));
291 local_agent.begin_discardable_simulation();
292
293 let dispatch = {
294 let root = runtime.root.as_mut().expect("parallel_uct root missing");
295 dispatch_rollout::<M>(
296 root,
297 &mut runtime.next_node_id,
298 local_agent.as_mut(),
299 horizon,
300 workers,
301 bu_uct_m_max,
302 )
303 };
304 pending.push(PendingRollout {
305 task_index,
306 agent: local_agent,
307 remaining_horizon: dispatch.remaining_horizon,
308 path: dispatch.path,
309 });
310 }
311
312 let completed = pending
313 .into_par_iter()
314 .map(|mut task| CompletedRollout {
315 task_index: task.task_index,
316 path: task.path,
317 tail_reward: random_rollout(task.agent.as_mut(), task.remaining_horizon),
318 })
319 .collect::<Vec<_>>();
320 let mut completed = completed;
321 completed.sort_by_key(|task| task.task_index);
322 complete_update_batch(runtime, &completed, gamma);
323
324 dispatched += batch_size;
325 }
326
327 if samples > 0 {
328 let root = runtime.root.as_ref().expect("parallel_uct root missing");
329 return best_action_after_positive_budget(root, agent);
330 }
331
332 best_action(runtime.root.as_ref(), agent)
333}
334
335fn planner_task_seed(planner_seed: u64, task_index: usize) -> u64 {
336 planner_seed ^ ((task_index as u64).wrapping_mul(PARALLEL_PLANNER_SEED_SALT))
337}
338
339fn best_action<M: ModeState>(
340 root: Option<&DecisionNode<M>>,
341 agent: &mut dyn AgentSimulator,
342) -> Action {
343 let Some(root) = root else {
344 return agent.gen_range(agent.get_num_actions().get()) as Action;
345 };
346 best_action_from_action_values(
347 root.action_edges
348 .iter()
349 .enumerate()
350 .filter_map(|(action_idx, edge)| {
351 edge.as_ref()
352 .filter(|edge| edge.completed_n() > 0)
353 .map(|edge| (action_idx, edge.completed_q()))
354 }),
355 agent.get_num_actions(),
356 agent,
357 )
358}
359
360fn best_action_after_positive_budget<M: ModeState>(
361 root: &DecisionNode<M>,
362 agent: &mut dyn AgentSimulator,
363) -> Action {
364 debug_assert!(
365 root.action_edges
366 .iter()
367 .filter_map(Option::as_ref)
368 .any(|edge| edge.completed_n() > 0),
369 "positive-budget parallel_uct search must leave at least one completed root edge"
370 );
371
372 best_action(Some(root), agent)
373}
374
375#[cfg(test)]
376fn root_has_completed_edge<M: ModeState>(root: &DecisionNode<M>) -> bool {
377 root.action_edges
378 .iter()
379 .filter_map(Option::as_ref)
380 .any(|edge| edge.completed_n() > 0)
381}
382
383fn prune_tree<M: ModeState>(
384 runtime: &mut ParallelRuntime<M>,
385 agent: &dyn AgentSimulator,
386 prev_obs_stream: &[PerceptVal],
387 prev_rew: Reward,
388 prev_act: Action,
389) {
390 let Some(mut old_root) = runtime.root.take() else {
391 runtime.root = Some(runtime.fresh_decision_node());
392 return;
393 };
394
395 let action_edge = old_root
396 .action_edges
397 .get_mut(prev_act as usize)
398 .and_then(Option::take);
399 let Some(mut action_edge) = action_edge else {
400 runtime.root = Some(runtime.fresh_decision_node());
401 return;
402 };
403
404 let key = prune_key(agent, prev_obs_stream, prev_rew);
405 runtime.root = action_edge
406 .chance_mut()
407 .percept_children
408 .remove(&key)
409 .or_else(|| Some(runtime.fresh_decision_node()));
410}
411
412fn bootstrap_root<M: ModeState>(
413 runtime: &mut ParallelRuntime<M>,
414 agent: &mut dyn AgentSimulator,
415 remaining_horizon: usize,
416 task_index: usize,
417) -> CompletedRollout {
418 debug_assert!(remaining_horizon > 0);
419 let root = runtime.root.as_mut().expect("parallel_uct root missing");
420 let action_idx = choose_bootstrap_action(root, agent);
421
422 agent.model_update_action(action_idx as Action);
423 let (observations, immediate_reward) = agent.gen_percepts_and_update();
424 let outcome = PerceptOutcome::new(observations, immediate_reward);
425
426 let parent_node_id = root.id;
427 let edge = root.action_edges[action_idx]
428 .as_mut()
429 .expect("parallel_uct bootstrap action edge missing");
430 if !edge.chance().percept_children.contains_key(&outcome) {
431 let id = runtime.next_node_id;
432 runtime.next_node_id += 1;
433 edge.chance_mut()
434 .percept_children
435 .insert(outcome.clone(), DecisionNode::new(id));
436 }
437 edge.on_incomplete_update();
438 let child_node_id = edge
439 .chance()
440 .percept_children
441 .get(&outcome)
442 .expect("parallel_uct bootstrap percept child missing")
443 .id;
444
445 let path = vec![ParallelPathStep {
446 parent_node_id,
447 action_idx,
448 child_node_id,
449 outcome,
450 }];
451 let tail_reward = random_rollout(agent, remaining_horizon.saturating_sub(1));
452 CompletedRollout {
453 task_index,
454 path,
455 tail_reward,
456 }
457}
458
459fn choose_bootstrap_action<M: ModeState>(
460 root: &mut DecisionNode<M>,
461 agent: &mut dyn AgentSimulator,
462) -> usize {
463 let num_actions = agent.get_num_actions();
464 ensure_action_slots(&mut root.action_edges, num_actions.get());
465
466 let mut unvisited = Vec::new();
467 for action_idx in 0..num_actions.get() {
468 match root.action_edges.get(action_idx).and_then(Option::as_ref) {
469 None => unvisited.push(action_idx),
470 Some(edge) if edge.completed_n() == 0 && edge.effective_visits() == 0 => {
471 unvisited.push(action_idx);
472 }
473 Some(_) => {}
474 }
475 }
476
477 let selected = unvisited[agent.gen_range(unvisited.len())];
478 if root.action_edges[selected].is_none() {
479 root.action_edges[selected] = Some(M::Edge::new());
480 }
481 selected
482}
483
484fn dispatch_rollout<M: ModeState>(
485 node: &mut DecisionNode<M>,
486 next_node_id: &mut u64,
487 agent: &mut dyn AgentSimulator,
488 remaining_horizon: usize,
489 workers: usize,
490 bu_uct_m_max: Option<f64>,
491) -> DispatchRollout {
492 let mut path = Vec::new();
493 let remaining_horizon = dispatch_rollout_into::<M>(
494 node,
495 next_node_id,
496 agent,
497 remaining_horizon,
498 workers,
499 bu_uct_m_max,
500 &mut path,
501 );
502 DispatchRollout {
503 remaining_horizon,
504 path,
505 }
506}
507
508fn dispatch_rollout_into<M: ModeState>(
509 node: &mut DecisionNode<M>,
510 next_node_id: &mut u64,
511 agent: &mut dyn AgentSimulator,
512 remaining_horizon: usize,
513 workers: usize,
514 bu_uct_m_max: Option<f64>,
515 path: &mut Vec<ParallelPathStep>,
516) -> usize {
517 if remaining_horizon == 0 || node.visits == 0 {
518 return remaining_horizon;
519 }
520
521 let num_actions = agent.get_num_actions();
522 ensure_action_slots(&mut node.action_edges, num_actions.get());
523
524 let action_idx = if let Some(unvisited) =
525 choose_uniform_unvisited(agent, &node.action_edges, num_actions.get())
526 {
527 node.action_edges[unvisited] = Some(M::Edge::new());
528 unvisited
529 } else {
530 let Some(action_idx) =
531 select_existing_action(node, agent, remaining_horizon, workers, bu_uct_m_max)
532 else {
533 return remaining_horizon;
534 };
535 action_idx
536 };
537
538 agent.model_update_action(action_idx as Action);
539 let (observations, immediate_reward) = agent.gen_percepts_and_update();
540 let outcome = PerceptOutcome::new(observations, immediate_reward);
541 let parent_node_id = node.id;
542
543 let edge = node.action_edges[action_idx]
544 .as_mut()
545 .expect("parallel_uct action edge missing");
546 if !edge.chance().percept_children.contains_key(&outcome) {
547 let id = *next_node_id;
548 *next_node_id += 1;
549 edge.chance_mut()
550 .percept_children
551 .insert(outcome.clone(), DecisionNode::new(id));
552 }
553 edge.on_incomplete_update();
554 let child = edge
555 .chance_mut()
556 .percept_children
557 .get_mut(&outcome)
558 .expect("parallel_uct percept child missing after insertion");
559 let child_node_id = child.id;
560
561 path.push(ParallelPathStep {
562 parent_node_id,
563 action_idx,
564 child_node_id,
565 outcome,
566 });
567
568 dispatch_rollout_into::<M>(
569 child,
570 next_node_id,
571 agent,
572 remaining_horizon - 1,
573 workers,
574 bu_uct_m_max,
575 path,
576 )
577}
578
579fn select_existing_action<M: ModeState>(
580 node: &DecisionNode<M>,
581 agent: &mut dyn AgentSimulator,
582 remaining_horizon: usize,
583 workers: usize,
584 bu_uct_m_max: Option<f64>,
585) -> Option<usize> {
586 let total_overline_n = node
587 .action_edges
588 .iter()
589 .filter_map(Option::as_ref)
590 .map(EdgeOps::effective_visits)
591 .sum::<u32>();
592 let log_total = ((total_overline_n.max(1)) as f64).ln().max(0.0);
593 let c = agent.get_explore_exploit_ratio().max(0.0);
594
595 let mut best_score = -f64::INFINITY;
596 let mut best_action = None;
597 let mut num_maximal_actions = 0usize;
598
599 for (action_idx, edge) in node.action_edges.iter().enumerate() {
600 let Some(edge) = edge.as_ref() else {
601 continue;
602 };
603 let overline_n = edge.effective_visits();
604 if overline_n == 0 || !M::edge_is_selectable(edge, workers, bu_uct_m_max) {
605 continue;
606 }
607
608 let normalized_value = agent.norm_reward_for_horizon(edge.completed_q(), remaining_horizon);
609 let exploration = c * ((2.0 * log_total) / (overline_n as f64)).sqrt();
610 let score = normalized_value + exploration;
611 debug_assert!(
612 score.is_finite(),
613 "parallel_uct UCB score must be finite for visited action edges"
614 );
615
616 match score.total_cmp(&best_score) {
617 std::cmp::Ordering::Greater => {
618 best_score = score;
619 best_action = Some(action_idx);
620 num_maximal_actions = 1;
621 }
622 std::cmp::Ordering::Equal => {
623 num_maximal_actions += 1;
624 if agent.gen_range(num_maximal_actions) == 0 {
625 best_action = Some(action_idx);
626 }
627 }
628 std::cmp::Ordering::Less => {}
629 }
630 }
631
632 best_action
633}
634
635#[cfg(test)]
636fn incomplete_update<M: ModeState>(runtime: &mut ParallelRuntime<M>, path: &[ParallelPathStep]) {
637 if path.is_empty() {
638 return;
639 }
640 let mut current = runtime.root.as_mut().expect("parallel_uct root missing");
641 for step in path {
642 let edge = current.action_edges[step.action_idx]
643 .as_mut()
644 .expect("parallel_uct action edge missing during incomplete_update");
645 edge.on_incomplete_update();
646 current = edge
647 .chance_mut()
648 .percept_children
649 .get_mut(&step.outcome)
650 .expect("parallel_uct percept child missing during incomplete_update");
651 }
652}
653
654fn complete_update_batch<M: ModeState>(
655 runtime: &mut ParallelRuntime<M>,
656 completed: &[CompletedRollout],
657 gamma: f64,
658) {
659 let mut epoch_state = M::EpochState::default();
660 for task in completed {
661 let root = runtime.root.as_mut().expect("parallel_uct root missing");
662 complete_update_node::<M>(
663 root,
664 &task.path,
665 0,
666 task.tail_reward,
667 gamma,
668 &mut epoch_state,
669 );
670 }
671}
672
673fn complete_update_node<M: ModeState>(
674 node: &mut DecisionNode<M>,
675 path: &[ParallelPathStep],
676 depth: usize,
677 tail_reward: f64,
678 gamma: f64,
679 epoch_state: &mut M::EpochState,
680) -> f64 {
681 node.visits += 1;
682 if depth == path.len() {
683 return tail_reward;
684 }
685
686 let step = &path[depth];
687 let edge = node.action_edges[step.action_idx]
688 .as_mut()
689 .expect("parallel_uct action edge missing during complete_update");
690 let child = edge
691 .chance_mut()
692 .percept_children
693 .get_mut(&step.outcome)
694 .expect("parallel_uct percept child missing during complete_update");
695 let downstream =
696 complete_update_node::<M>(child, path, depth + 1, tail_reward, gamma, epoch_state);
697
698 let reward = (step.outcome.reward() as f64) + gamma * downstream;
699 M::complete_edge(
700 edge,
701 BuEpochKey {
702 parent_node_id: step.parent_node_id,
703 action_idx: step.action_idx,
704 child_node_id: step.child_node_id,
705 },
706 reward,
707 epoch_state,
708 );
709 reward
710}
711
712enum ParallelPlannerState {
713 Wu(ParallelRuntime<WuMode>),
714 Bu(ParallelRuntime<BuMode>),
715}
716
717struct ParallelRuntime<M: ModeState> {
718 root: Option<DecisionNode<M>>,
719 next_node_id: u64,
720}
721
722impl<M: ModeState> ParallelRuntime<M> {
723 fn new() -> Self {
724 Self {
725 root: Some(DecisionNode::new(0)),
726 next_node_id: 1,
727 }
728 }
729
730 fn fresh_decision_node(&mut self) -> DecisionNode<M> {
731 let id = self.next_node_id;
732 self.next_node_id += 1;
733 DecisionNode::new(id)
734 }
735}
736
737trait ModeState: Copy {
738 type Edge: EdgeOps<Self>;
739 type EpochState: Default;
740
741 fn edge_is_selectable(edge: &Self::Edge, workers: usize, bu_uct_m_max: Option<f64>) -> bool;
742
743 fn complete_edge(
744 edge: &mut Self::Edge,
745 key: BuEpochKey,
746 reward: f64,
747 epoch_state: &mut Self::EpochState,
748 );
749}
750
751trait EdgeOps<M: ModeState>: Clone {
752 fn new() -> Self;
753 fn chance(&self) -> &ChanceNode<M>;
754 fn chance_mut(&mut self) -> &mut ChanceNode<M>;
755 fn effective_visits(&self) -> u32;
756 fn completed_q(&self) -> f64;
757 fn completed_n(&self) -> u32;
758 fn on_incomplete_update(&mut self);
759}
760
761#[derive(Clone, Copy)]
762struct WuMode;
763
764#[derive(Clone, Copy)]
765struct BuMode;
766
767#[derive(Clone)]
768struct DecisionNode<M: ModeState> {
769 id: u64,
770 visits: u32,
771 action_edges: Vec<Option<M::Edge>>,
772}
773
774impl<M: ModeState> DecisionNode<M> {
775 fn new(id: u64) -> Self {
776 Self {
777 id,
778 visits: 0,
779 action_edges: Vec::new(),
780 }
781 }
782}
783
784#[derive(Clone)]
785struct ChanceNode<M: ModeState> {
786 percept_children: PerceptMap<DecisionNode<M>>,
787}
788
789impl<M: ModeState> Default for ChanceNode<M> {
790 fn default() -> Self {
791 Self {
792 percept_children: PerceptMap::default(),
793 }
794 }
795}
796
797#[derive(Clone)]
798struct WuActionEdge {
799 q: f64,
800 n: u32,
801 o: u32,
802 child: ChanceNode<WuMode>,
803}
804
805impl EdgeOps<WuMode> for WuActionEdge {
806 fn new() -> Self {
807 Self {
808 q: 0.0,
809 n: 0,
810 o: 0,
811 child: ChanceNode::default(),
812 }
813 }
814
815 fn chance(&self) -> &ChanceNode<WuMode> {
816 &self.child
817 }
818
819 fn chance_mut(&mut self) -> &mut ChanceNode<WuMode> {
820 &mut self.child
821 }
822
823 fn effective_visits(&self) -> u32 {
824 self.n + self.o
825 }
826
827 fn completed_q(&self) -> f64 {
828 self.q
829 }
830
831 fn completed_n(&self) -> u32 {
832 self.n
833 }
834
835 fn on_incomplete_update(&mut self) {
836 self.o += 1;
837 }
838}
839
840impl ModeState for WuMode {
841 type Edge = WuActionEdge;
842 type EpochState = ();
843
844 fn edge_is_selectable(_edge: &Self::Edge, _workers: usize, _bu_uct_m_max: Option<f64>) -> bool {
845 true
846 }
847
848 fn complete_edge(
849 edge: &mut Self::Edge,
850 _key: BuEpochKey,
851 reward: f64,
852 _epoch_state: &mut Self::EpochState,
853 ) {
854 edge.o = edge.o.saturating_sub(1);
855 edge.q = (reward + (edge.n as f64) * edge.q) / ((edge.n + 1) as f64);
856 edge.n += 1;
857 }
858}
859
860#[derive(Clone)]
861struct BuActionEdge {
862 q: f64,
863 n: u32,
864 o: u32,
865 o_bar: f64,
868 child: ChanceNode<BuMode>,
869}
870
871impl EdgeOps<BuMode> for BuActionEdge {
872 fn new() -> Self {
873 Self {
874 q: 0.0,
875 n: 0,
876 o: 0,
877 o_bar: 0.0,
878 child: ChanceNode::default(),
879 }
880 }
881
882 fn chance(&self) -> &ChanceNode<BuMode> {
883 &self.child
884 }
885
886 fn chance_mut(&mut self) -> &mut ChanceNode<BuMode> {
887 &mut self.child
888 }
889
890 fn effective_visits(&self) -> u32 {
891 self.n + self.o
892 }
893
894 fn completed_q(&self) -> f64 {
895 self.q
896 }
897
898 fn completed_n(&self) -> u32 {
899 self.n
900 }
901
902 fn on_incomplete_update(&mut self) {
903 self.o += 1;
904 let overline_n = self.effective_visits();
905 if overline_n > 0 {
906 self.o_bar =
907 (((overline_n - 1) as f64) * self.o_bar + (self.o as f64)) / (overline_n as f64);
908 }
909 }
910}
911
912impl ModeState for BuMode {
913 type Edge = BuActionEdge;
914 type EpochState = BuEpochState;
915
916 fn edge_is_selectable(edge: &Self::Edge, workers: usize, bu_uct_m_max: Option<f64>) -> bool {
917 let m_max = bu_uct_m_max.expect("BU mode requires bu_uct_m_max");
918 edge.o_bar < m_max * (workers as f64)
919 }
920
921 fn complete_edge(
922 edge: &mut Self::Edge,
923 key: BuEpochKey,
924 reward: f64,
925 epoch_state: &mut Self::EpochState,
926 ) {
927 edge.o = edge.o.saturating_sub(1);
928 edge.update_bu_epoch(key, reward, epoch_state);
933 }
934}
935
936impl BuActionEdge {
937 fn update_bu_epoch(&mut self, key: BuEpochKey, reward: f64, epoch_state: &mut BuEpochState) {
938 use std::collections::hash_map::Entry;
939
940 match epoch_state.groups.entry(key) {
941 Entry::Vacant(entry) => {
942 entry.insert(BuGroupStat {
943 mean: reward,
944 count: 1,
945 });
946 self.q = if self.n == 0 {
947 reward
948 } else {
949 (((self.n as f64) * self.q) + reward) / ((self.n + 1) as f64)
950 };
951 self.n += 1;
952 }
953 Entry::Occupied(mut entry) => {
954 let old_mean = entry.get().mean;
955 let stat = entry.get_mut();
956 stat.count += 1;
957 stat.mean = old_mean + (reward - old_mean) / (stat.count as f64);
958 if self.n > 0 {
959 self.q += (stat.mean - old_mean) / (self.n as f64);
960 }
961 }
962 }
963 }
964}
965
966#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
967struct BuEpochKey {
968 parent_node_id: u64,
969 action_idx: usize,
970 child_node_id: u64,
971}
972
973#[derive(Clone, Copy)]
974struct BuGroupStat {
975 mean: f64,
976 count: u32,
977}
978
979#[derive(Default)]
980struct BuEpochState {
981 groups: HashMap<BuEpochKey, BuGroupStat>,
982}
983
984struct DispatchRollout {
985 remaining_horizon: usize,
986 path: Vec<ParallelPathStep>,
987}
988
989struct PendingRollout {
990 task_index: usize,
991 agent: Box<dyn AgentSimulator>,
992 remaining_horizon: usize,
993 path: Vec<ParallelPathStep>,
994}
995
996struct CompletedRollout {
997 task_index: usize,
998 path: Vec<ParallelPathStep>,
999 tail_reward: f64,
1000}
1001
1002#[derive(Clone)]
1003struct ParallelPathStep {
1004 parent_node_id: u64,
1005 action_idx: usize,
1006 child_node_id: u64,
1007 outcome: PerceptOutcome,
1008}
1009
1010#[cfg(test)]
1011type WuDecisionNode = DecisionNode<WuMode>;
1012#[cfg(test)]
1013type BuDecisionNode = DecisionNode<BuMode>;
1014#[cfg(test)]
1015type BuChanceNode = ChanceNode<BuMode>;
1016
1017#[cfg(test)]
1018impl ParallelUctPlanner {
1019 fn wu_root(&self) -> &WuDecisionNode {
1020 match &self.state {
1021 ParallelPlannerState::Wu(runtime) => runtime.root.as_ref().expect("WU root"),
1022 ParallelPlannerState::Bu(_) => panic!("expected WU planner"),
1023 }
1024 }
1025
1026 fn bu_root(&self) -> &BuDecisionNode {
1027 match &self.state {
1028 ParallelPlannerState::Bu(runtime) => runtime.root.as_ref().expect("BU root"),
1029 ParallelPlannerState::Wu(_) => panic!("expected BU planner"),
1030 }
1031 }
1032
1033 fn set_wu_root(&mut self, root: WuDecisionNode) {
1034 match &mut self.state {
1035 ParallelPlannerState::Wu(runtime) => runtime.root = Some(root),
1036 ParallelPlannerState::Bu(_) => panic!("expected WU planner"),
1037 }
1038 }
1039
1040 fn set_bu_root(&mut self, root: BuDecisionNode) {
1041 match &mut self.state {
1042 ParallelPlannerState::Bu(runtime) => runtime.root = Some(root),
1043 ParallelPlannerState::Wu(_) => panic!("expected BU planner"),
1044 }
1045 }
1046}
1047
1048#[cfg(test)]
1049mod tests {
1050 use super::*;
1051 use crate::aixi::mcts::RhoUctPlanner;
1052 use std::sync::atomic::{AtomicUsize, Ordering};
1053
1054 fn workers(n: usize) -> NonZeroUsize {
1060 NonZeroUsize::new(n).expect("test fixtures must use non-zero worker counts")
1061 }
1062
1063 #[derive(Clone)]
1064 struct DeterministicRewardAgent {
1065 last_action: Action,
1066 emit_reward: bool,
1067 }
1068
1069 impl AgentSimulator for DeterministicRewardAgent {
1070 fn get_num_actions(&self) -> ActionAlphabet {
1071 ActionAlphabet::try_from_usize(2).expect("test fixture action alphabet must be valid")
1072 }
1073
1074 fn get_num_observation_bits(&self) -> usize {
1075 1
1076 }
1077
1078 fn get_num_reward_bits(&self) -> usize {
1079 1
1080 }
1081
1082 fn horizon(&self) -> usize {
1083 1
1084 }
1085
1086 fn max_reward(&self) -> Reward {
1087 1
1088 }
1089
1090 fn min_reward(&self) -> Reward {
1091 0
1092 }
1093
1094 fn get_explore_exploit_ratio(&self) -> f64 {
1095 0.0
1096 }
1097
1098 fn model_update_action(&mut self, action: Action) {
1099 self.last_action = action;
1100 self.emit_reward = false;
1101 }
1102
1103 fn gen_percept_and_update(&mut self, _bits: usize) -> u64 {
1104 if self.emit_reward {
1105 self.emit_reward = false;
1106 self.last_action
1107 } else {
1108 self.emit_reward = true;
1109 0
1110 }
1111 }
1112
1113 fn model_revert(&mut self, _steps: usize) {
1114 self.emit_reward = false;
1115 }
1116
1117 fn gen_range(&mut self, _end: usize) -> usize {
1118 0
1119 }
1120
1121 fn gen_f64(&mut self) -> f64 {
1122 0.0
1123 }
1124
1125 fn boxed_clone_with_seed(&self, seed: u64) -> Box<dyn AgentSimulator> {
1126 let _ = seed;
1127 Box::new(self.clone())
1128 }
1129 }
1130
1131 #[derive(Clone)]
1132 struct ThresholdProbeAgent {
1133 num_actions: usize,
1134 model_updates: usize,
1135 range_result: usize,
1136 }
1137
1138 impl AgentSimulator for ThresholdProbeAgent {
1139 fn get_num_actions(&self) -> ActionAlphabet {
1140 ActionAlphabet::try_from_usize(self.num_actions)
1141 .expect("test fixture action alphabet must be valid")
1142 }
1143
1144 fn get_num_observation_bits(&self) -> usize {
1145 1
1146 }
1147
1148 fn get_num_reward_bits(&self) -> usize {
1149 1
1150 }
1151
1152 fn horizon(&self) -> usize {
1153 1
1154 }
1155
1156 fn max_reward(&self) -> Reward {
1157 1
1158 }
1159
1160 fn min_reward(&self) -> Reward {
1161 0
1162 }
1163
1164 fn get_explore_exploit_ratio(&self) -> f64 {
1165 0.0
1166 }
1167
1168 fn model_update_action(&mut self, _action: Action) {
1169 self.model_updates += 1;
1170 }
1171
1172 fn gen_percept_and_update(&mut self, _bits: usize) -> u64 {
1173 0
1174 }
1175
1176 fn model_revert(&mut self, _steps: usize) {}
1177
1178 fn gen_range(&mut self, end: usize) -> usize {
1179 self.range_result.min(end.saturating_sub(1))
1180 }
1181
1182 fn gen_f64(&mut self) -> f64 {
1183 0.0
1184 }
1185
1186 fn boxed_clone_with_seed(&self, _seed: u64) -> Box<dyn AgentSimulator> {
1187 Box::new(self.clone())
1188 }
1189 }
1190
1191 #[derive(Clone)]
1192 struct CounterAgent {
1193 clone_count: Arc<AtomicUsize>,
1194 begin_count: Arc<AtomicUsize>,
1195 discardable_begin_count: Arc<AtomicUsize>,
1196 model_updates: Arc<AtomicUsize>,
1197 planning_horizon: usize,
1198 last_action: Action,
1199 emit_reward: bool,
1200 }
1201
1202 impl CounterAgent {
1203 fn new_with_horizon(planning_horizon: usize) -> Self {
1204 Self {
1205 clone_count: Arc::new(AtomicUsize::new(0)),
1206 begin_count: Arc::new(AtomicUsize::new(0)),
1207 discardable_begin_count: Arc::new(AtomicUsize::new(0)),
1208 model_updates: Arc::new(AtomicUsize::new(0)),
1209 planning_horizon,
1210 last_action: 0,
1211 emit_reward: false,
1212 }
1213 }
1214 }
1215
1216 impl AgentSimulator for CounterAgent {
1217 fn get_num_actions(&self) -> ActionAlphabet {
1218 ActionAlphabet::try_from_usize(2).expect("test fixture action alphabet must be valid")
1219 }
1220
1221 fn get_num_observation_bits(&self) -> usize {
1222 1
1223 }
1224
1225 fn get_num_reward_bits(&self) -> usize {
1226 1
1227 }
1228
1229 fn horizon(&self) -> usize {
1230 self.planning_horizon
1231 }
1232
1233 fn max_reward(&self) -> Reward {
1234 1
1235 }
1236
1237 fn min_reward(&self) -> Reward {
1238 0
1239 }
1240
1241 fn get_explore_exploit_ratio(&self) -> f64 {
1242 0.0
1243 }
1244
1245 fn begin_simulation(&mut self) {
1246 self.begin_count.fetch_add(1, Ordering::SeqCst);
1247 }
1248
1249 fn begin_discardable_simulation(&mut self) {
1250 self.discardable_begin_count.fetch_add(1, Ordering::SeqCst);
1251 }
1252
1253 fn model_update_action(&mut self, action: Action) {
1254 self.last_action = action;
1255 self.emit_reward = false;
1256 self.model_updates.fetch_add(1, Ordering::SeqCst);
1257 }
1258
1259 fn gen_percept_and_update(&mut self, _bits: usize) -> u64 {
1260 if self.emit_reward {
1261 self.emit_reward = false;
1262 self.last_action
1263 } else {
1264 self.emit_reward = true;
1265 0
1266 }
1267 }
1268
1269 fn model_revert(&mut self, _steps: usize) {
1270 self.emit_reward = false;
1271 }
1272
1273 fn gen_range(&mut self, _end: usize) -> usize {
1274 0
1275 }
1276
1277 fn gen_f64(&mut self) -> f64 {
1278 0.0
1279 }
1280
1281 fn boxed_clone_with_seed(&self, _seed: u64) -> Box<dyn AgentSimulator> {
1282 self.clone_count.fetch_add(1, Ordering::SeqCst);
1283 Box::new(self.clone())
1284 }
1285 }
1286
1287 #[derive(Clone)]
1288 struct SeedRecordingAgent {
1289 recorded_seeds: Arc<Mutex<Vec<u64>>>,
1290 last_action: Action,
1291 emit_reward: bool,
1292 clone_seed: u64,
1293 }
1294
1295 impl SeedRecordingAgent {
1296 fn new() -> Self {
1297 Self {
1298 recorded_seeds: Arc::new(Mutex::new(Vec::new())),
1299 last_action: 0,
1300 emit_reward: false,
1301 clone_seed: 0,
1302 }
1303 }
1304 }
1305
1306 impl AgentSimulator for SeedRecordingAgent {
1307 fn get_num_actions(&self) -> ActionAlphabet {
1308 ActionAlphabet::try_from_usize(2).expect("test fixture action alphabet must be valid")
1309 }
1310
1311 fn get_num_observation_bits(&self) -> usize {
1312 1
1313 }
1314
1315 fn get_num_reward_bits(&self) -> usize {
1316 1
1317 }
1318
1319 fn horizon(&self) -> usize {
1320 1
1321 }
1322
1323 fn max_reward(&self) -> Reward {
1324 1
1325 }
1326
1327 fn min_reward(&self) -> Reward {
1328 0
1329 }
1330
1331 fn get_explore_exploit_ratio(&self) -> f64 {
1332 0.0
1333 }
1334
1335 fn model_update_action(&mut self, action: Action) {
1336 self.last_action = action;
1337 self.emit_reward = false;
1338 }
1339
1340 fn gen_percept_and_update(&mut self, _bits: usize) -> u64 {
1341 if self.emit_reward {
1342 self.emit_reward = false;
1343 (self.clone_seed ^ self.last_action) & 1
1344 } else {
1345 self.emit_reward = true;
1346 0
1347 }
1348 }
1349
1350 fn model_revert(&mut self, _steps: usize) {
1351 self.emit_reward = false;
1352 }
1353
1354 fn gen_range(&mut self, _end: usize) -> usize {
1355 0
1356 }
1357
1358 fn gen_f64(&mut self) -> f64 {
1359 0.25
1360 }
1361
1362 fn boxed_clone_with_seed(&self, seed: u64) -> Box<dyn AgentSimulator> {
1363 self.recorded_seeds.lock().expect("seed list").push(seed);
1364 Box::new(Self {
1365 recorded_seeds: Arc::clone(&self.recorded_seeds),
1366 last_action: 0,
1367 emit_reward: false,
1368 clone_seed: seed,
1369 })
1370 }
1371 }
1372
1373 fn step(
1374 parent_node_id: u64,
1375 action_idx: usize,
1376 child_node_id: u64,
1377 observation: u64,
1378 reward: Reward,
1379 ) -> ParallelPathStep {
1380 ParallelPathStep {
1381 parent_node_id,
1382 action_idx,
1383 child_node_id,
1384 outcome: PerceptOutcome::new(vec![observation], reward),
1385 }
1386 }
1387
1388 fn two_step_wu_root() -> WuDecisionNode {
1389 WuDecisionNode {
1390 id: 10,
1391 visits: 2,
1392 action_edges: vec![
1393 Some(WuActionEdge {
1394 q: 1.0,
1395 n: 1,
1396 o: 0,
1397 child: ChanceNode {
1398 percept_children: HashMap::from([(
1399 PerceptOutcome::new(vec![0], 0),
1400 WuDecisionNode {
1401 id: 11,
1402 visits: 2,
1403 action_edges: vec![
1404 Some(WuActionEdge {
1405 q: 1.0,
1406 n: 1,
1407 o: 0,
1408 child: ChanceNode {
1409 percept_children: HashMap::from([(
1410 PerceptOutcome::new(vec![0], 0),
1411 WuDecisionNode::new(12),
1412 )]),
1413 },
1414 }),
1415 Some(WuActionEdge {
1416 q: 0.0,
1417 n: 1,
1418 o: 0,
1419 child: ChanceNode::default(),
1420 }),
1421 ],
1422 },
1423 )]),
1424 },
1425 }),
1426 Some(WuActionEdge {
1427 q: 0.0,
1428 n: 1,
1429 o: 0,
1430 child: ChanceNode::default(),
1431 }),
1432 ],
1433 }
1434 }
1435
1436 fn retained_parent_for(next_root: WuDecisionNode) -> WuDecisionNode {
1437 WuDecisionNode {
1438 id: 9,
1439 visits: 1,
1440 action_edges: vec![Some(WuActionEdge {
1441 q: 0.0,
1442 n: 1,
1443 o: 0,
1444 child: ChanceNode {
1445 percept_children: HashMap::from([(PerceptOutcome::new(vec![0], 0), next_root)]),
1446 },
1447 })],
1448 }
1449 }
1450
1451 #[test]
1457 fn planner_new_rejects_invalid_bu_threshold_at_api_boundary() {
1458 for invalid in [0.0, 1.0, -0.1, 1.1] {
1459 let err = match ParallelUctPlanner::new(workers(2), Some(invalid)) {
1460 Ok(_) => panic!("invalid BU threshold must be rejected"),
1461 Err(err) => err,
1462 };
1463 assert_eq!(err, ParallelUctPlannerInitError::InvalidBuUctMMax);
1464 }
1465 }
1466
1467 #[test]
1468 fn wu_workers_one_matches_rho_uct_on_deterministic_agent() {
1469 let mut seq_agent = DeterministicRewardAgent {
1470 last_action: 0,
1471 emit_reward: false,
1472 };
1473 let mut par_agent = seq_agent.clone();
1474
1475 let mut sequential = RhoUctPlanner::new();
1476 let mut parallel =
1477 ParallelUctPlanner::new(workers(1), None).expect("valid parallel_uct planner");
1478
1479 let seq_action = sequential.search(&mut seq_agent, &[0], 0, 0, 16);
1480 let par_action = parallel
1481 .search(&mut par_agent, &[0], 0, 0, 16)
1482 .expect("positive-horizon parallel_uct search");
1483
1484 assert_eq!(seq_action, par_action);
1485 assert_eq!(parallel.wu_root().visits, 16);
1486 }
1487
1488 #[test]
1489 fn fresh_root_bootstrap_produces_completed_root_edge_for_positive_budget() {
1490 let mut agent = DeterministicRewardAgent {
1491 last_action: 0,
1492 emit_reward: false,
1493 };
1494 let mut planner =
1495 ParallelUctPlanner::new(workers(4), None).expect("valid parallel_uct planner");
1496 let action = planner
1497 .search(&mut agent, &[0], 0, 0, 4)
1498 .expect("positive-horizon parallel_uct search");
1499
1500 let root = planner.wu_root();
1501 assert!(action < 2);
1502 assert_eq!(root.visits, 4);
1503 assert!(
1504 root_has_completed_edge(root),
1505 "positive-budget search on a fresh retained root must complete at least one root edge"
1506 );
1507 assert!(
1508 root.action_edges
1509 .iter()
1510 .filter_map(Option::as_ref)
1511 .all(|edge| edge.o == 0),
1512 "all incomplete counts must be cleared after the search batch completes"
1513 );
1514 }
1515
1516 #[test]
1517 fn zero_sample_budget_does_not_bootstrap_or_clone_even_at_zero_horizon() {
1518 let mut agent = CounterAgent::new_with_horizon(0);
1519 let mut planner =
1520 ParallelUctPlanner::new(workers(4), None).expect("valid parallel_uct planner");
1521
1522 let action = planner
1523 .search(&mut agent, &[0], 0, 0, 0)
1524 .expect("zero-sample parallel_uct search should not require positive horizon");
1525 assert_eq!(action, 0);
1526 assert_eq!(agent.clone_count.load(Ordering::SeqCst), 0);
1527 assert_eq!(agent.begin_count.load(Ordering::SeqCst), 0);
1528 assert_eq!(agent.discardable_begin_count.load(Ordering::SeqCst), 0);
1529 assert_eq!(agent.model_updates.load(Ordering::SeqCst), 0);
1530 let root = planner.wu_root();
1531 assert_eq!(root.visits, 0);
1532 assert!(
1533 !root_has_completed_edge(root),
1534 "zero-budget search must not synthesize completed root edges"
1535 );
1536 }
1537
1538 #[test]
1539 fn positive_budget_search_rejects_zero_horizon_at_api_boundary() {
1540 let mut agent = CounterAgent::new_with_horizon(0);
1541 let mut planner =
1542 ParallelUctPlanner::new(workers(4), None).expect("valid parallel_uct planner");
1543
1544 let err = planner
1545 .search(&mut agent, &[0], 0, 0, 1)
1546 .expect_err("positive-budget zero-horizon search must be rejected");
1547 assert_eq!(
1548 err,
1549 ParallelUctSearchError::PositiveSamplesRequirePositiveHorizon
1550 );
1551 assert_eq!(agent.clone_count.load(Ordering::SeqCst), 0);
1552 assert_eq!(agent.begin_count.load(Ordering::SeqCst), 0);
1553 assert_eq!(agent.discardable_begin_count.load(Ordering::SeqCst), 0);
1554 assert_eq!(agent.model_updates.load(Ordering::SeqCst), 0);
1555 let root = planner.wu_root();
1556 assert_eq!(root.visits, 0);
1557 assert!(
1558 !root_has_completed_edge(root),
1559 "rejected zero-horizon search must leave the retained root untouched"
1560 );
1561 }
1562
1563 #[test]
1564 fn single_sample_bootstrap_avoids_empty_root_fallback() {
1565 let mut agent = DeterministicRewardAgent {
1566 last_action: 0,
1567 emit_reward: false,
1568 };
1569 let mut planner =
1570 ParallelUctPlanner::new(workers(4), None).expect("valid parallel_uct planner");
1571
1572 let action = planner
1573 .search(&mut agent, &[0], 0, 0, 1)
1574 .expect("positive-horizon parallel_uct search");
1575 assert_eq!(action, 0);
1576 let root = planner.wu_root();
1577 assert_eq!(root.visits, 1);
1578 assert!(
1579 root_has_completed_edge(root),
1580 "single-sample retained-root bootstrap must leave one completed root edge"
1581 );
1582 }
1583
1584 #[test]
1585 fn positive_budget_uses_discardable_simulation_hook_for_clones() {
1586 let mut agent = CounterAgent::new_with_horizon(2);
1587 let mut planner =
1588 ParallelUctPlanner::new(workers(2), None).expect("valid parallel_uct planner");
1589
1590 let _action = planner
1591 .search(&mut agent, &[0], 0, 0, 5)
1592 .expect("positive-horizon parallel_uct search");
1593
1594 assert_eq!(agent.clone_count.load(Ordering::SeqCst), 5);
1595 assert_eq!(agent.discardable_begin_count.load(Ordering::SeqCst), 5);
1596 assert_eq!(
1597 agent.begin_count.load(Ordering::SeqCst),
1598 0,
1599 "parallel_uct cloned rollouts must not open reversible simulation scopes",
1600 );
1601 }
1602
1603 #[test]
1604 fn retained_root_same_percept_branch_reuses_bootstrapped_subtree_skeleton() {
1605 let mut agent = DeterministicRewardAgent {
1606 last_action: 0,
1607 emit_reward: false,
1608 };
1609 let mut planner =
1610 ParallelUctPlanner::new(workers(2), None).expect("valid parallel_uct planner");
1611 let action = planner
1612 .search(&mut agent, &[0], 0, 0, 1)
1613 .expect("positive-horizon parallel_uct search");
1614 let root = planner.wu_root();
1615 let edge = root.action_edges[action as usize]
1616 .as_ref()
1617 .expect("root edge");
1618 let retained = edge
1619 .chance()
1620 .percept_children
1621 .get(&PerceptOutcome::new(vec![0], 0))
1622 .expect("retained subtree");
1623 let retained_id = retained.id;
1624
1625 let _follow_up = planner
1626 .search(&mut agent, &[0], 0, action, 0)
1627 .expect("zero-sample parallel_uct search should not require positive horizon");
1628 let next_root = planner.wu_root();
1629 assert_eq!(next_root.id, retained_id);
1630 assert_eq!(next_root.visits, 1);
1631 }
1632
1633 #[test]
1634 fn prune_to_fresh_root_bootstrap_preserves_positive_budget_invariant() {
1635 let mut agent = DeterministicRewardAgent {
1636 last_action: 0,
1637 emit_reward: false,
1638 };
1639 let mut planner =
1640 ParallelUctPlanner::new(workers(2), None).expect("valid parallel_uct planner");
1641 planner.set_wu_root(WuDecisionNode {
1642 id: 0,
1643 visits: 8,
1644 action_edges: vec![Some(WuActionEdge {
1645 q: 0.0,
1646 n: 1,
1647 o: 0,
1648 child: ChanceNode {
1649 percept_children: HashMap::from([(
1650 PerceptOutcome::new(vec![0], 0),
1651 WuDecisionNode::new(1),
1652 )]),
1653 },
1654 })],
1655 });
1656
1657 let second_action = planner
1658 .search(&mut agent, &[0], 0, 0, 1)
1659 .expect("positive-horizon parallel_uct search");
1660 assert_eq!(second_action, 0);
1661 let root = planner.wu_root();
1662 assert_eq!(root.visits, 1);
1663 assert!(
1664 root_has_completed_edge(root),
1665 "prune-to-fresh-root with positive budget must still complete a root edge"
1666 );
1667 }
1668
1669 #[test]
1670 fn bootstrap_and_batched_rollouts_use_absolute_task_indices_for_seeding() {
1671 let mut agent = SeedRecordingAgent::new();
1672 let mut planner =
1673 ParallelUctPlanner::new(workers(4), None).expect("valid parallel_uct planner");
1674
1675 let _action = planner
1676 .search(&mut agent, &[0], 0, 0, 5)
1677 .expect("positive-horizon parallel_uct search");
1678
1679 let planner_seed = 0.25f64.to_bits();
1680 let expected = (0..5)
1681 .map(|task_index| planner_task_seed(planner_seed, task_index))
1682 .collect::<Vec<_>>();
1683 let seen = agent.recorded_seeds.lock().expect("seed list").clone();
1684 assert_eq!(seen, expected);
1685 }
1686
1687 #[test]
1688 fn dispatch_rollout_records_root_to_leaf_path_and_updates_each_edge_once() {
1689 let mut node = two_step_wu_root();
1690 let mut next_node_id = 13u64;
1691 let mut agent = DeterministicRewardAgent {
1692 last_action: 0,
1693 emit_reward: false,
1694 };
1695
1696 let dispatch =
1697 dispatch_rollout::<WuMode>(&mut node, &mut next_node_id, &mut agent, 3, 1, None);
1698 assert_eq!(dispatch.remaining_horizon, 1);
1699 assert_eq!(dispatch.path.len(), 2);
1700 assert_eq!(dispatch.path[0].parent_node_id, 10);
1701 assert_eq!(dispatch.path[0].action_idx, 0);
1702 assert_eq!(dispatch.path[0].child_node_id, 11);
1703 assert_eq!(dispatch.path[0].outcome, PerceptOutcome::new(vec![0], 0));
1704 assert_eq!(dispatch.path[1].parent_node_id, 11);
1705 assert_eq!(dispatch.path[1].action_idx, 0);
1706 assert_eq!(dispatch.path[1].child_node_id, 12);
1707 assert_eq!(dispatch.path[1].outcome, PerceptOutcome::new(vec![0], 0));
1708
1709 let root_edge = node.action_edges[0].as_ref().expect("root edge");
1710 let child = root_edge
1711 .chance()
1712 .percept_children
1713 .get(&PerceptOutcome::new(vec![0], 0))
1714 .expect("child node");
1715 let child_edge = child.action_edges[0].as_ref().expect("child edge");
1716 assert_eq!(root_edge.o, 1);
1717 assert_eq!(child_edge.o, 1);
1718 }
1719
1720 #[test]
1721 fn ordinary_dispatch_path_completes_single_rollout_without_residual_incomplete_counts() {
1722 let mut agent = CounterAgent::new_with_horizon(3);
1723 let mut planner =
1724 ParallelUctPlanner::new(workers(1), None).expect("valid parallel_uct planner");
1725 planner.set_wu_root(retained_parent_for(two_step_wu_root()));
1726
1727 let action = planner
1728 .search(&mut agent, &[0], 0, 0, 1)
1729 .expect("positive-horizon parallel_uct search");
1730 assert_eq!(action, 0);
1731
1732 let root = planner.wu_root();
1733 assert_eq!(root.id, 10);
1734 let root_edge = root.action_edges[0].as_ref().expect("root edge");
1735 let child = root_edge
1736 .chance()
1737 .percept_children
1738 .get(&PerceptOutcome::new(vec![0], 0))
1739 .expect("child node");
1740 let child_edge = child.action_edges[0].as_ref().expect("child edge");
1741 assert_eq!(root_edge.o, 0);
1742 assert_eq!(child_edge.o, 0);
1743 assert_eq!(root_edge.n, 2);
1744 assert_eq!(child_edge.n, 2);
1745 }
1746
1747 #[test]
1748 fn bu_thresholding_skips_oversubscribed_edges() {
1749 let mut agent = DeterministicRewardAgent {
1750 last_action: 0,
1751 emit_reward: false,
1752 };
1753 let mut node = BuDecisionNode::new(0);
1754 node.visits = 8;
1755 node.action_edges = vec![
1756 Some(BuActionEdge {
1757 q: 0.9,
1758 n: 4,
1759 o: 0,
1760 o_bar: 2.0,
1761 child: BuChanceNode::default(),
1762 }),
1763 Some(BuActionEdge {
1764 q: 0.1,
1765 n: 4,
1766 o: 0,
1767 o_bar: 0.0,
1768 child: BuChanceNode::default(),
1769 }),
1770 ];
1771
1772 let selected = select_existing_action::<BuMode>(&node, &mut agent, 1, 2, Some(0.5));
1773 assert_eq!(
1774 selected,
1775 Some(1),
1776 "BU-UCT should skip edges whose average incomplete count exceeds the threshold"
1777 );
1778 }
1779
1780 #[test]
1781 fn bu_thresholding_returns_none_when_all_edges_are_oversubscribed() {
1782 let mut agent = DeterministicRewardAgent {
1783 last_action: 0,
1784 emit_reward: false,
1785 };
1786 let mut node = BuDecisionNode::new(0);
1787 node.visits = 8;
1788 node.action_edges = vec![
1789 Some(BuActionEdge {
1790 q: 0.9,
1791 n: 4,
1792 o: 0,
1793 o_bar: 2.0,
1794 child: BuChanceNode::default(),
1795 }),
1796 Some(BuActionEdge {
1797 q: 0.1,
1798 n: 4,
1799 o: 0,
1800 o_bar: 2.0,
1801 child: BuChanceNode::default(),
1802 }),
1803 ];
1804
1805 let selected = select_existing_action::<BuMode>(&node, &mut agent, 1, 2, Some(0.5));
1806 assert_eq!(selected, None);
1807 }
1808
1809 #[test]
1810 fn bu_thresholding_stops_at_current_node_when_all_expanded_children_forbidden() {
1811 let mut node = BuDecisionNode::new(0);
1812 node.visits = 8;
1813 node.action_edges = vec![
1814 Some(BuActionEdge {
1815 q: 0.9,
1816 n: 4,
1817 o: 0,
1818 o_bar: 2.0,
1819 child: BuChanceNode::default(),
1820 }),
1821 Some(BuActionEdge {
1822 q: 0.1,
1823 n: 4,
1824 o: 0,
1825 o_bar: 2.0,
1826 child: BuChanceNode::default(),
1827 }),
1828 ];
1829 let mut next_node_id = 1u64;
1830 let mut agent = ThresholdProbeAgent {
1831 num_actions: 2,
1832 model_updates: 0,
1833 range_result: 0,
1834 };
1835
1836 let dispatch =
1837 dispatch_rollout::<BuMode>(&mut node, &mut next_node_id, &mut agent, 3, 2, Some(0.5));
1838 assert!(
1839 dispatch.path.is_empty(),
1840 "when every expanded child is threshold-forbidden, BU-core must stop at the current node"
1841 );
1842 assert_eq!(dispatch.remaining_horizon, 3);
1843 assert_eq!(
1844 agent.model_updates, 0,
1845 "threshold stop must not traverse or expand a forbidden edge"
1846 );
1847 }
1848
1849 #[test]
1850 fn bu_threshold_stop_expands_unexpanded_legal_action_synchronously() {
1851 let mut node = BuDecisionNode::new(0);
1852 node.visits = 8;
1853 node.action_edges = vec![
1854 Some(BuActionEdge {
1855 q: 0.9,
1856 n: 4,
1857 o: 0,
1858 o_bar: 2.0,
1859 child: BuChanceNode::default(),
1860 }),
1861 Some(BuActionEdge {
1862 q: 0.1,
1863 n: 4,
1864 o: 0,
1865 o_bar: 2.0,
1866 child: BuChanceNode::default(),
1867 }),
1868 None,
1869 ];
1870 let mut next_node_id = 7u64;
1871 let mut agent = ThresholdProbeAgent {
1872 num_actions: 3,
1873 model_updates: 0,
1874 range_result: 0,
1875 };
1876
1877 let dispatch =
1878 dispatch_rollout::<BuMode>(&mut node, &mut next_node_id, &mut agent, 3, 2, Some(0.5));
1879 assert_eq!(dispatch.path.len(), 1);
1880 assert_eq!(dispatch.path[0].action_idx, 2);
1881 assert_eq!(dispatch.path[0].parent_node_id, 0);
1882 assert_eq!(dispatch.path[0].child_node_id, 7);
1883 assert_eq!(dispatch.remaining_horizon, 2);
1884 assert_eq!(agent.model_updates, 1);
1885 assert!(
1886 node.action_edges[2].is_some(),
1887 "BU-core must synchronously expand an unexpanded legal action after threshold stop"
1888 );
1889 let edge = node.action_edges[2].as_ref().expect("expanded edge");
1890 let child = edge
1891 .chance()
1892 .percept_children
1893 .get(&PerceptOutcome::new(vec![0], 0))
1894 .expect("expanded child");
1895 assert_eq!(edge.o, 1);
1896 assert_eq!(child.id, dispatch.path[0].child_node_id);
1897 }
1898
1899 #[test]
1900 fn bu_thresholding_uses_configured_worker_budget_not_batch_size() {
1901 let mut planner =
1902 ParallelUctPlanner::new(workers(4), Some(0.8)).expect("valid parallel_uct planner");
1903 let retained_root = BuDecisionNode {
1904 id: 1,
1905 visits: 8,
1906 action_edges: vec![Some(BuActionEdge {
1907 q: 0.2,
1908 n: 1,
1909 o: 0,
1910 o_bar: 1.0,
1911 child: BuChanceNode::default(),
1912 })],
1913 };
1914 planner.set_bu_root(BuDecisionNode {
1915 id: 0,
1916 visits: 8,
1917 action_edges: vec![Some(BuActionEdge {
1918 n: 1,
1919 q: 0.0,
1920 o: 0,
1921 o_bar: 0.0,
1922 child: BuChanceNode {
1923 percept_children: HashMap::from([(
1924 PerceptOutcome::new(vec![0], 0),
1925 retained_root,
1926 )]),
1927 },
1928 })],
1929 });
1930 let mut agent = ThresholdProbeAgent {
1931 num_actions: 1,
1932 model_updates: 0,
1933 range_result: 0,
1934 };
1935
1936 let action = planner
1937 .search(&mut agent, &[0], 0, 0, 1)
1938 .expect("positive-horizon parallel_uct search");
1939 assert_eq!(action, 0);
1940 let root = planner.bu_root();
1941 let edge = root.action_edges[0].as_ref().expect("root edge");
1942 assert_eq!(
1943 edge.n, 2,
1944 "with configured workers=4 and m_max=0.8, O_bar=1.0 must remain admissible even for samples=1"
1945 );
1946 }
1947
1948 #[test]
1949 fn bu_grouped_backpropagation_keeps_group_count_constant_for_same_child_origin() {
1950 let mut edge = BuActionEdge::new();
1951 let mut epoch_state = BuEpochState::default();
1952 let key = BuEpochKey {
1953 parent_node_id: 0,
1954 action_idx: 0,
1955 child_node_id: 7,
1956 };
1957
1958 edge.update_bu_epoch(key, 1.0, &mut epoch_state);
1959 assert_eq!(edge.n, 1);
1960 assert_eq!(edge.q, 1.0);
1961
1962 edge.update_bu_epoch(key, 0.0, &mut epoch_state);
1963 assert_eq!(edge.n, 1);
1964 assert!((edge.q - 0.5).abs() < 1e-12);
1965 }
1966
1967 #[test]
1968 fn bu_batch_updates_q_once_from_same_origin_group_mean() {
1969 let mut planner =
1970 ParallelUctPlanner::new(workers(4), Some(0.8)).expect("valid parallel_uct planner");
1971 planner.set_bu_root(BuDecisionNode {
1972 id: 0,
1973 visits: 0,
1974 action_edges: vec![Some(BuActionEdge {
1975 q: 0.0,
1976 n: 0,
1977 o: 2,
1978 o_bar: 0.0,
1979 child: BuChanceNode {
1980 percept_children: HashMap::from([(
1981 PerceptOutcome::new(vec![0], 0),
1982 BuDecisionNode::new(1),
1983 )]),
1984 },
1985 })],
1986 });
1987
1988 let completed = vec![
1989 CompletedRollout {
1990 task_index: 1,
1991 path: vec![step(0, 0, 1, 0, 0)],
1992 tail_reward: 1.0,
1993 },
1994 CompletedRollout {
1995 task_index: 0,
1996 path: vec![step(0, 0, 1, 0, 0)],
1997 tail_reward: 0.0,
1998 },
1999 ];
2000 match &mut planner.state {
2001 ParallelPlannerState::Bu(runtime) => complete_update_batch(runtime, &completed, 1.0),
2002 ParallelPlannerState::Wu(_) => panic!("expected BU planner"),
2003 }
2004
2005 let root = planner.bu_root();
2006 let edge = root.action_edges[0].as_ref().expect("root edge");
2007 assert_eq!(root.visits, 2);
2008 assert_eq!(edge.n, 1);
2009 assert_eq!(edge.o, 0);
2010 assert!((edge.q - 0.5).abs() < 1e-12);
2011 }
2012
2013 #[test]
2014 fn bu_batch_distinguishes_different_child_origins_at_same_ancestor_edge() {
2015 let mut planner =
2016 ParallelUctPlanner::new(workers(4), Some(0.8)).expect("valid parallel_uct planner");
2017 planner.set_bu_root(BuDecisionNode {
2018 id: 0,
2019 visits: 0,
2020 action_edges: vec![Some(BuActionEdge {
2021 q: 0.0,
2022 n: 0,
2023 o: 2,
2024 o_bar: 0.0,
2025 child: BuChanceNode {
2026 percept_children: HashMap::from([
2027 (PerceptOutcome::new(vec![0], 0), BuDecisionNode::new(1)),
2028 (PerceptOutcome::new(vec![1], 0), BuDecisionNode::new(2)),
2029 ]),
2030 },
2031 })],
2032 });
2033
2034 let completed = vec![
2035 CompletedRollout {
2036 task_index: 0,
2037 path: vec![step(0, 0, 1, 0, 0)],
2038 tail_reward: 1.0,
2039 },
2040 CompletedRollout {
2041 task_index: 1,
2042 path: vec![step(0, 0, 2, 1, 0)],
2043 tail_reward: 0.0,
2044 },
2045 ];
2046 match &mut planner.state {
2047 ParallelPlannerState::Bu(runtime) => complete_update_batch(runtime, &completed, 1.0),
2048 ParallelPlannerState::Wu(_) => panic!("expected BU planner"),
2049 }
2050
2051 let root = planner.bu_root();
2052 let edge = root.action_edges[0].as_ref().expect("root edge");
2053 assert_eq!(edge.n, 2);
2054 assert!((edge.q - 0.5).abs() < 1e-12);
2055 }
2056
2057 #[test]
2058 fn bu_grouping_is_local_to_each_ancestor_edge() {
2059 let mut planner =
2060 ParallelUctPlanner::new(workers(4), Some(0.8)).expect("valid parallel_uct planner");
2061 planner.set_bu_root(BuDecisionNode {
2062 id: 0,
2063 visits: 0,
2064 action_edges: vec![Some(BuActionEdge {
2065 q: 0.0,
2066 n: 0,
2067 o: 2,
2068 o_bar: 0.0,
2069 child: BuChanceNode {
2070 percept_children: HashMap::from([(
2071 PerceptOutcome::new(vec![0], 0),
2072 BuDecisionNode {
2073 id: 1,
2074 visits: 0,
2075 action_edges: vec![Some(BuActionEdge {
2076 q: 0.0,
2077 n: 0,
2078 o: 2,
2079 o_bar: 0.0,
2080 child: BuChanceNode {
2081 percept_children: HashMap::from([
2082 (PerceptOutcome::new(vec![10], 0), BuDecisionNode::new(2)),
2083 (PerceptOutcome::new(vec![11], 0), BuDecisionNode::new(3)),
2084 ]),
2085 },
2086 })],
2087 },
2088 )]),
2089 },
2090 })],
2091 });
2092
2093 let completed = vec![
2094 CompletedRollout {
2095 task_index: 0,
2096 path: vec![step(0, 0, 1, 0, 0), step(1, 0, 2, 10, 0)],
2097 tail_reward: 1.0,
2098 },
2099 CompletedRollout {
2100 task_index: 1,
2101 path: vec![step(0, 0, 1, 0, 0), step(1, 0, 3, 11, 0)],
2102 tail_reward: 0.0,
2103 },
2104 ];
2105 match &mut planner.state {
2106 ParallelPlannerState::Bu(runtime) => complete_update_batch(runtime, &completed, 1.0),
2107 ParallelPlannerState::Wu(_) => panic!("expected BU planner"),
2108 }
2109
2110 let root = planner.bu_root();
2111 let root_edge = root.action_edges[0].as_ref().expect("root edge");
2112 let child = root_edge
2113 .chance()
2114 .percept_children
2115 .get(&PerceptOutcome::new(vec![0], 0))
2116 .expect("child node");
2117 let child_edge = child.action_edges[0].as_ref().expect("child edge");
2118
2119 assert_eq!(root_edge.n, 1);
2120 assert_eq!(child_edge.n, 2);
2121 }
2122
2123 #[test]
2124 fn incomplete_update_tracks_o_bar_recurrence() {
2125 let mut planner =
2126 ParallelUctPlanner::new(workers(4), Some(0.8)).expect("valid parallel_uct planner");
2127 planner.set_bu_root(BuDecisionNode {
2128 id: 0,
2129 visits: 0,
2130 action_edges: vec![Some(BuActionEdge {
2131 q: 0.0,
2132 n: 0,
2133 o: 0,
2134 o_bar: 0.0,
2135 child: BuChanceNode {
2136 percept_children: HashMap::from([(
2137 PerceptOutcome::new(vec![0], 0),
2138 BuDecisionNode::new(1),
2139 )]),
2140 },
2141 })],
2142 });
2143
2144 let path = vec![step(0, 0, 1, 0, 0)];
2145 match &mut planner.state {
2146 ParallelPlannerState::Bu(runtime) => incomplete_update(runtime, &path),
2147 ParallelPlannerState::Wu(_) => panic!("expected BU planner"),
2148 }
2149 let root = planner.bu_root();
2150 let edge = root.action_edges[0].as_ref().expect("edge");
2151 assert_eq!(edge.o, 1);
2152 assert!((edge.o_bar - 1.0).abs() < 1e-12);
2153
2154 match &mut planner.state {
2155 ParallelPlannerState::Bu(runtime) => incomplete_update(runtime, &path),
2156 ParallelPlannerState::Wu(_) => panic!("expected BU planner"),
2157 }
2158 let root = planner.bu_root();
2159 let edge = root.action_edges[0].as_ref().expect("edge");
2160 assert_eq!(edge.o, 2);
2161 assert!((edge.o_bar - 1.5).abs() < 1e-12);
2162 }
2163
2164 #[test]
2165 fn bu_thresholding_uses_o_bar_not_live_o_after_completion() {
2166 let mut planner =
2167 ParallelUctPlanner::new(workers(4), Some(0.5)).expect("valid parallel_uct planner");
2168 planner.set_bu_root(BuDecisionNode {
2169 id: 0,
2170 visits: 8,
2171 action_edges: vec![
2172 Some(BuActionEdge {
2173 q: 0.9,
2174 n: 4,
2175 o: 1,
2176 o_bar: 2.0,
2177 child: BuChanceNode {
2178 percept_children: HashMap::from([(
2179 PerceptOutcome::new(vec![0], 0),
2180 BuDecisionNode::new(1),
2181 )]),
2182 },
2183 }),
2184 Some(BuActionEdge {
2185 q: 0.1,
2186 n: 4,
2187 o: 0,
2188 o_bar: 0.0,
2189 child: BuChanceNode::default(),
2190 }),
2191 ],
2192 });
2193
2194 let completed = vec![CompletedRollout {
2195 task_index: 0,
2196 path: vec![step(0, 0, 1, 0, 0)],
2197 tail_reward: 0.0,
2198 }];
2199 match &mut planner.state {
2200 ParallelPlannerState::Bu(runtime) => complete_update_batch(runtime, &completed, 1.0),
2201 ParallelPlannerState::Wu(_) => panic!("expected BU planner"),
2202 }
2203
2204 let root = planner.bu_root();
2205 let forbidden = root.action_edges[0].as_ref().expect("forbidden edge");
2206 assert_eq!(
2207 forbidden.o, 0,
2208 "completion must clear the live incomplete count"
2209 );
2210 assert!(
2211 (forbidden.o_bar - 2.0).abs() < 1e-12,
2212 "BU-core Part 1 keeps the paper-style one-sided o_bar statistic after completion"
2213 );
2214
2215 let mut agent = DeterministicRewardAgent {
2216 last_action: 0,
2217 emit_reward: false,
2218 };
2219 let selected = select_existing_action::<BuMode>(root, &mut agent, 1, 4, Some(0.5));
2220 assert_eq!(
2221 selected,
2222 Some(1),
2223 "post-completion BU thresholding must continue to follow o_bar rather than current live o"
2224 );
2225 }
2226
2227 #[test]
2228 fn bu_bootstrap_uses_singleton_epoch_completion() {
2229 let mut agent = DeterministicRewardAgent {
2230 last_action: 0,
2231 emit_reward: false,
2232 };
2233 let mut planner =
2234 ParallelUctPlanner::new(workers(4), Some(0.5)).expect("valid parallel_uct planner");
2235
2236 let action = planner
2237 .search(&mut agent, &[0], 0, 0, 1)
2238 .expect("positive-horizon parallel_uct search");
2239 assert_eq!(action, 0);
2240 let root = planner.bu_root();
2241 assert_eq!(root.visits, 1);
2242 let edge = root.action_edges[0].as_ref().expect("bootstrap edge");
2243 assert_eq!(edge.n, 1);
2244 assert_eq!(edge.o, 0);
2245 }
2246
2247 #[test]
2248 fn final_root_choice_uses_completed_q_not_bu_thresholding() {
2249 let mut planner =
2250 ParallelUctPlanner::new(workers(2), Some(0.5)).expect("valid parallel_uct planner");
2251 planner.set_bu_root(BuDecisionNode {
2252 id: 0,
2253 visits: 8,
2254 action_edges: vec![
2255 Some(BuActionEdge {
2256 q: 0.9,
2257 n: 4,
2258 o: 0,
2259 o_bar: 2.0,
2260 child: BuChanceNode::default(),
2261 }),
2262 Some(BuActionEdge {
2263 q: 0.1,
2264 n: 4,
2265 o: 0,
2266 o_bar: 0.0,
2267 child: BuChanceNode::default(),
2268 }),
2269 ],
2270 });
2271
2272 let mut agent = DeterministicRewardAgent {
2273 last_action: 0,
2274 emit_reward: false,
2275 };
2276 let action = best_action::<BuMode>(Some(planner.bu_root()), &mut agent);
2277 assert_eq!(action, 0);
2278 }
2279
2280 #[test]
2281 fn parallel_search_is_thread_count_independent() {
2282 fn run_with_threads(threads: usize) -> (Action, u32, Vec<(u32, u32)>) {
2283 let pool = rayon::ThreadPoolBuilder::new()
2284 .num_threads(threads)
2285 .build()
2286 .expect("thread pool");
2287 pool.install(|| {
2288 let mut agent = DeterministicRewardAgent {
2289 last_action: 0,
2290 emit_reward: false,
2291 };
2292 let mut planner = ParallelUctPlanner::new(workers(4), Some(0.8))
2293 .expect("valid parallel_uct planner");
2294 let action = planner
2295 .search(&mut agent, &[0], 0, 0, 32)
2296 .expect("positive-horizon parallel_uct search");
2297 let root = planner.bu_root();
2298 let child_stats = root
2299 .action_edges
2300 .iter()
2301 .map(|edge| edge.as_ref().map_or((0, 0), |edge| (edge.n, edge.o)))
2302 .collect::<Vec<_>>();
2303 (action, root.visits, child_stats)
2304 })
2305 }
2306
2307 let one = run_with_threads(1);
2308 let two = run_with_threads(2);
2309 let four = run_with_threads(4);
2310
2311 assert_eq!(one, two);
2312 assert_eq!(two, four);
2313 }
2314}