1use super::*;
3
4#[derive(Clone, Debug, Default)]
5pub struct SequentialScheduler;
6
7#[derive(Clone, Debug)]
8pub struct ParallelScheduler {
9 max_workers: usize,
10}
11
12impl ParallelScheduler {
13 pub fn new(max_workers: usize) -> Result<Self> {
14 if max_workers == 0 {
15 return Err(DagMlError::RuntimeValidation(
16 "parallel scheduler max_workers must be at least 1".to_string(),
17 ));
18 }
19 Ok(Self { max_workers })
20 }
21
22 pub fn max_workers(&self) -> usize {
23 self.max_workers
24 }
25}
26
27#[derive(Clone, Debug)]
28pub(crate) struct PhaseScope {
29 pub(crate) phase: Phase,
30 pub(crate) variant_id: Option<VariantId>,
31 pub(crate) variant: Option<VariantExecutionSpec>,
32 pub(crate) fold_id: Option<FoldId>,
33 pub(crate) seed_root: Option<u64>,
34}
35
36#[derive(Clone, Debug)]
37pub(crate) struct ReplayPredictionCacheContract {
38 pub(crate) requirement: BundlePredictionRequirement,
39 pub(crate) cache: BundlePredictionCacheRecord,
40}
41
42pub(crate) struct MaterializedReplayArtifacts {
43 pub(crate) handles: BTreeMap<NodeId, BTreeMap<String, HandleRef>>,
44 pub(crate) inputs: BTreeMap<NodeId, BTreeMap<String, ArtifactInputSpec>>,
45}
46
47fn prediction_output_ports_for_node(plan: &ExecutionPlan, node_id: &NodeId) -> Result<Vec<String>> {
48 let node = plan
49 .graph_plan
50 .graph
51 .nodes
52 .iter()
53 .find(|node| node.id == *node_id)
54 .ok_or_else(|| {
55 DagMlError::RuntimeValidation(format!(
56 "node `{node_id}` is absent from the execution graph"
57 ))
58 })?;
59 let mut ports = node
60 .ports
61 .outputs
62 .iter()
63 .filter(|port| port.kind == PortKind::Prediction)
64 .map(|port| port.name.clone())
65 .collect::<Vec<_>>();
66 ports.sort();
67 Ok(ports)
68}
69
70fn normalize_prediction_result_port(
71 node_id: &NodeId,
72 block_kind: &str,
73 producer_port: &mut Option<String>,
74 prediction_ports: &[String],
75) -> Result<()> {
76 if let Some(port) = producer_port.as_ref() {
77 if port.trim().is_empty() {
78 return Err(DagMlError::RuntimeValidation(format!(
79 "node `{node_id}` emitted {block_kind} with blank producer_port"
80 )));
81 }
82 if !prediction_ports.iter().any(|candidate| candidate == port) {
83 return Err(DagMlError::RuntimeValidation(format!(
84 "node `{node_id}` emitted {block_kind} for undeclared or non-prediction output port `{port}`; declared prediction ports are {:?}",
85 prediction_ports
86 )));
87 }
88 return Ok(());
89 }
90 match prediction_ports {
91 [only] => {
92 *producer_port = Some(only.clone());
93 Ok(())
94 }
95 [] => Err(DagMlError::RuntimeValidation(format!(
96 "node `{node_id}` emitted {block_kind} without producer_port but declares no prediction output port"
97 ))),
98 _ => Err(DagMlError::RuntimeValidation(format!(
99 "node `{node_id}` emitted {block_kind} without producer_port but declares {} prediction output ports {:?}; multi-output controllers must emit producer_port explicitly",
100 prediction_ports.len(),
101 prediction_ports
102 ))),
103 }
104}
105
106pub(crate) fn normalize_result_prediction_ports(
107 plan: &ExecutionPlan,
108 task: &NodeTask,
109 result: &mut NodeResult,
110) -> Result<()> {
111 if result.predictions.is_empty()
112 && result.observation_predictions.is_empty()
113 && result.aggregated_predictions.is_empty()
114 && result.explanations.is_empty()
115 {
116 return Ok(());
117 }
118 let prediction_ports = prediction_output_ports_for_node(plan, &task.node_plan.node_id)?;
119 for block in &mut result.predictions {
120 normalize_prediction_result_port(
121 &task.node_plan.node_id,
122 "prediction block",
123 &mut block.producer_port,
124 &prediction_ports,
125 )?;
126 }
127 for block in &mut result.observation_predictions {
128 normalize_prediction_result_port(
129 &task.node_plan.node_id,
130 "observation prediction block",
131 &mut block.producer_port,
132 &prediction_ports,
133 )?;
134 }
135 for block in &mut result.aggregated_predictions {
136 normalize_prediction_result_port(
137 &task.node_plan.node_id,
138 "aggregated prediction block",
139 &mut block.producer_port,
140 &prediction_ports,
141 )?;
142 }
143 for block in &mut result.explanations {
144 normalize_prediction_result_port(
145 &task.node_plan.node_id,
146 "explanation block",
147 &mut block.producer_port,
148 &prediction_ports,
149 )?;
150 }
151 Ok(())
152}
153
154#[derive(Default)]
155pub(crate) struct PhaseScopeResources<'a> {
156 pub(crate) data_provider: Option<&'a dyn RuntimeDataProvider>,
157 pub(crate) replay_artifact_handles: Option<&'a BTreeMap<NodeId, BTreeMap<String, HandleRef>>>,
158 pub(crate) replay_artifact_inputs:
159 Option<&'a BTreeMap<NodeId, BTreeMap<String, ArtifactInputSpec>>>,
160 pub(crate) replay_bundle_id: Option<&'a BundleId>,
161 pub(crate) data_envelopes: Option<&'a BTreeMap<String, ExternalDataPlanEnvelope>>,
162 pub(crate) prediction_cache_store: Option<&'a dyn RuntimePredictionCacheStore>,
163 pub(crate) prediction_cache_contracts:
164 Option<&'a BTreeMap<String, ReplayPredictionCacheContract>>,
165 pub(crate) artifact_store: Option<&'a mut InMemoryArtifactStore>,
166}
167
168impl SequentialScheduler {
169 pub fn execute_hpo_campaign(
176 &self,
177 plan: &ExecutionPlan,
178 controllers: &RuntimeControllerRegistry,
179 data_provider: &dyn RuntimeDataProvider,
180 ctx: &RunContext,
181 hpo: &RuntimeHpoExecutionContext,
182 ) -> Result<RuntimeHpoCampaignResult> {
183 plan.validate()?;
184 hpo.validate_for_plan(plan)?;
185 let controller = controllers.get(&hpo.controller_id).ok_or_else(|| {
186 DagMlError::RuntimeValidation(format!(
187 "runtime HPO campaign controller `{}` is not registered",
188 hpo.controller_id
189 ))
190 })?;
191 let task = RuntimeHpoCampaignTask {
192 run_id: ctx.run_id.clone(),
193 operation_id: hpo.operation_id.clone(),
194 controller_id: hpo.controller_id.clone(),
195 target_node_id: hpo.target_node_id.clone(),
196 seed: ctx.root_seed,
197 };
198 let mut session = controller.create_tuner_session(&task, hpo)?;
199 let history_at_start = session.trial_history_len()?;
200 if history_at_start > hpo.trial_budget_total {
201 return Err(DagMlError::RuntimeValidation(format!(
202 "runtime HPO restored native history ({history_at_start}) exceeds total trial budget ({})",
203 hpo.trial_budget_total
204 )));
205 }
206 let remaining_trials = hpo.trial_budget_total - history_at_start;
207 let mut candidates = Vec::new();
208 let mut proposed_variant_ids = BTreeSet::new();
209 let mut trial_variants = BTreeMap::new();
214 let mut incumbent_variants = hpo.resume_variants.clone();
215 let mut terminal_trials = BTreeMap::new();
216 let mut completed_proposals = Vec::new();
217 let mut completed_reports = Vec::new();
218
219 for _ in 0..remaining_trials {
220 let Some(proposal) = session.ask()? else {
221 break;
222 };
223 if trial_variants
224 .insert(proposal.trial_id, proposal.variant.variant_id.clone())
225 .is_some()
226 {
227 return Err(DagMlError::RuntimeValidation(format!(
228 "runtime HPO session proposed duplicate trial `{}`",
229 proposal.trial_id
230 )));
231 }
232 if incumbent_variants
233 .insert(proposal.trial_id, proposal.variant.variant_id.clone())
234 .is_some()
235 {
236 return Err(DagMlError::RuntimeValidation(format!(
237 "runtime HPO session reused restored trial `{}`",
238 proposal.trial_id
239 )));
240 }
241 if !proposed_variant_ids.insert(proposal.variant.variant_id.clone()) {
242 return Err(DagMlError::RuntimeValidation(format!(
243 "runtime HPO session proposed duplicate variant `{}`",
244 proposal.variant.variant_id
245 )));
246 }
247 let mut candidate_plan = plan.clone();
248 candidate_plan.variants = vec![proposal.variant.clone()];
249 candidate_plan.validate()?;
250 let mut candidate_ctx =
251 RunContext::new(ctx.run_id.clone(), proposal.variant.seed.or(ctx.root_seed));
252 candidate_ctx.variant_id = Some(proposal.variant.variant_id.clone());
253
254 let evaluation = self.execute_hpo_candidate_fit_cv(
255 &candidate_plan,
256 controllers,
257 data_provider,
258 &mut candidate_ctx,
259 );
260 if let Err(error) = evaluation {
261 session.tell(
262 proposal.trial_id,
263 RuntimeHpoTerminal::Failed {
264 failure: RuntimeHpoFailure {
265 code: "DAGML_CV_ERROR".to_string(),
266 message: error.to_string(),
267 retryable: false,
268 },
269 },
270 )?;
271 terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Failed);
272 continue;
273 }
274 if let Err(error) = candidate_ctx
275 .collect_cross_fold_validation_scores(plan_oof_partition_mode(&candidate_plan))
276 {
277 session.tell(
278 proposal.trial_id,
279 RuntimeHpoTerminal::Failed {
280 failure: RuntimeHpoFailure {
281 code: "DAGML_SCORE_ERROR".to_string(),
282 message: error.to_string(),
283 retryable: false,
284 },
285 },
286 )?;
287 terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Failed);
288 continue;
289 }
290 let report = candidate_ctx
291 .score_collector
292 .iter()
293 .find(|report| {
294 report.producer_node == hpo.selection.producer_node
295 && report.producer_port.as_deref()
296 == Some(hpo.selection.producer_port.as_str())
297 && report.partition == PredictionPartition::Validation
298 && report
299 .fold_id
300 .as_ref()
301 .is_some_and(|fold| fold.as_str() == "avg")
302 })
303 .cloned();
304 let Some(mut report) = report else {
305 session.tell(
306 proposal.trial_id,
307 RuntimeHpoTerminal::Failed {
308 failure: RuntimeHpoFailure {
309 code: "DAGML_SCORE_MISSING".to_string(),
310 message: format!(
311 "runtime HPO trial `{}` emitted no target OOF average",
312 proposal.trial_id
313 ),
314 retryable: false,
315 },
316 },
317 )?;
318 terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Failed);
319 continue;
320 };
321 report.variant_id = Some(proposal.variant.variant_id.clone());
322 let score = report
323 .metrics
324 .get(hpo.selection.metric.name())
325 .copied()
326 .filter(|score| score.is_finite());
327 let Some(score) = score else {
328 session.tell(
329 proposal.trial_id,
330 RuntimeHpoTerminal::Failed {
331 failure: RuntimeHpoFailure {
332 code: "DAGML_SCORE_NONFINITE".to_string(),
333 message: format!(
334 "runtime HPO trial `{}` emitted no finite `{}` score",
335 proposal.trial_id,
336 hpo.selection.metric.name()
337 ),
338 retryable: false,
339 },
340 },
341 )?;
342 terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Failed);
343 continue;
344 };
345 let intermediate = RuntimeHpoIntermediate {
346 trial_id: proposal.trial_id,
347 step: 0,
348 score,
349 };
350 if session.report_intermediate(intermediate)? == RuntimeHpoIntermediateOutcome::Pruned {
351 terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Pruned);
352 continue;
353 }
354 session.tell(proposal.trial_id, RuntimeHpoTerminal::Completed { score })?;
355 terminal_trials.insert(proposal.trial_id, HpoTrialTerminalState::Completed);
356 completed_proposals.push(proposal.clone());
357 completed_reports.push(RuntimeHpoCompletedReport {
358 trial_id: proposal.trial_id,
359 variant_id: proposal.variant.variant_id.clone(),
360 report: report.clone(),
361 });
362
363 let mut validation_reports = candidate_ctx
364 .score_collector
365 .iter()
366 .filter(|item| item.partition == PredictionPartition::Validation)
367 .cloned()
368 .collect::<Vec<_>>();
369 for item in &mut validation_reports {
370 item.variant_id = Some(proposal.variant.variant_id.clone());
371 }
372 candidates.push(RuntimeHpoCandidateEvaluation {
373 validation_predictions: capture_variant_validation_predictions(
374 &proposal.variant.variant_id,
375 None,
376 &candidate_ctx,
377 ),
378 lineage: candidate_ctx.lineage.records().cloned().collect(),
379 proposal,
380 score,
381 validation_reports,
382 });
383 }
384
385 let history_at_checkpoint = session.trial_history_len()?;
386 if history_at_checkpoint != hpo.trial_budget_total {
387 return Err(DagMlError::RuntimeValidation(format!(
388 "runtime HPO native history ended at {history_at_checkpoint}, expected total trial budget {}",
389 hpo.trial_budget_total
390 )));
391 }
392
393 let checkpoint = RuntimeHpoCheckpointResult {
394 artifact: session.checkpoint()?,
395 provenance: hpo.provenance.clone(),
396 operation_id: hpo.operation_id.clone(),
397 controller_id: hpo.controller_id.clone(),
398 target_node_id: hpo.target_node_id.clone(),
399 completed_proposals,
400 completed_reports,
401 trial_history_len: history_at_checkpoint,
402 };
403 validate_hpo_checkpoint_result(
404 &checkpoint,
405 hpo,
406 &trial_variants,
407 &terminal_trials,
408 history_at_start,
409 )?;
410 let incumbent = session.incumbent(&incumbent_variants)?.ok_or_else(|| {
411 DagMlError::RuntimeValidation(
412 "native HPO campaign has no completed native incumbent after terminalization"
413 .to_string(),
414 )
415 })?;
416 if incumbent.metric != hpo.selection.metric.name()
417 || incumbent.direction != hpo.selection.direction
418 || incumbent_variants.get(&incumbent.trial_id) != Some(&incumbent.variant_id)
419 || !incumbent.score.is_finite()
420 {
421 return Err(DagMlError::RuntimeValidation(
422 "native HPO incumbent is not bound to this scheduler campaign's metric, direction, trial, and variant"
423 .to_string(),
424 ));
425 }
426 let terminal_trials = session.terminal_trial_snapshots(&incumbent_variants)?;
427 if terminal_trials.len() != history_at_checkpoint as usize
428 || terminal_trials
429 .windows(2)
430 .any(|pair| pair[0].trial.id >= pair[1].trial.id)
431 {
432 return Err(DagMlError::RuntimeValidation(
433 "native HPO terminal ledger is not a complete strictly ordered history".to_string(),
434 ));
435 }
436 Ok(RuntimeHpoCampaignResult {
437 operation_id: hpo.operation_id.clone(),
438 controller_id: hpo.controller_id.clone(),
439 target_node_id: hpo.target_node_id.clone(),
440 candidates,
441 checkpoint,
442 incumbent,
443 terminal_trials,
444 })
445 }
446
447 fn execute_hpo_candidate_fit_cv(
448 &self,
449 plan: &ExecutionPlan,
450 controllers: &RuntimeControllerRegistry,
451 data_provider: &dyn RuntimeDataProvider,
452 ctx: &mut RunContext,
453 ) -> Result<Vec<NodeResult>> {
454 let candidate_plan = plan;
455 let fold_ids = candidate_plan
456 .fold_set
457 .as_ref()
458 .map(|fold_set| {
459 fold_set
460 .folds
461 .iter()
462 .map(|fold| Some(fold.fold_id.clone()))
463 .collect::<Vec<_>>()
464 })
465 .unwrap_or_else(|| vec![None]);
466 let variant = candidate_plan
467 .variants
468 .first()
469 .expect("candidate plan has exactly one variant");
470 let mut results = Vec::new();
471 for fold_id in fold_ids {
472 results.extend(self.execute_phase_scope(
473 candidate_plan,
474 controllers,
475 ctx,
476 PhaseScope {
477 phase: Phase::FitCv,
478 variant_id: Some(variant.variant_id.clone()),
479 variant: Some(VariantExecutionSpec::from_plan(variant)),
480 fold_id,
481 seed_root: variant.seed.or(ctx.root_seed),
482 },
483 PhaseScopeResources {
484 data_provider: Some(data_provider),
485 ..Default::default()
486 },
487 )?);
488 }
489 Ok(results)
490 }
491
492 pub fn execute_phase(
493 &self,
494 plan: &ExecutionPlan,
495 controllers: &RuntimeControllerRegistry,
496 ctx: &mut RunContext,
497 phase: Phase,
498 ) -> Result<Vec<NodeResult>> {
499 plan.validate()?;
500 let variant_id = ctx.variant_id.clone();
501 let seed_root = ctx.root_seed;
502 self.execute_phase_scope(
503 plan,
504 controllers,
505 ctx,
506 PhaseScope {
507 phase,
508 variant_id,
509 variant: None,
510 fold_id: None,
511 seed_root,
512 },
513 PhaseScopeResources::default(),
514 )
515 }
516
517 pub fn execute_phase_with_data_provider(
518 &self,
519 plan: &ExecutionPlan,
520 controllers: &RuntimeControllerRegistry,
521 data_provider: &dyn RuntimeDataProvider,
522 ctx: &mut RunContext,
523 phase: Phase,
524 ) -> Result<Vec<NodeResult>> {
525 plan.validate()?;
526 let variant_id = ctx.variant_id.clone();
527 let seed_root = ctx.root_seed;
528 self.execute_phase_scope(
529 plan,
530 controllers,
531 ctx,
532 PhaseScope {
533 phase,
534 variant_id,
535 variant: None,
536 fold_id: None,
537 seed_root,
538 },
539 PhaseScopeResources {
540 data_provider: Some(data_provider),
541 ..Default::default()
542 },
543 )
544 }
545
546 pub fn execute_campaign_phase(
547 &self,
548 plan: &ExecutionPlan,
549 controllers: &RuntimeControllerRegistry,
550 ctx: &mut RunContext,
551 phase: Phase,
552 ) -> Result<Vec<NodeResult>> {
553 plan.validate()?;
554 let mut results = Vec::new();
555 let fold_ids = if phase == Phase::FitCv {
556 plan.fold_set
557 .as_ref()
558 .map(|fold_set| {
559 fold_set
560 .folds
561 .iter()
562 .map(|fold| Some(fold.fold_id.clone()))
563 .collect::<Vec<_>>()
564 })
565 .unwrap_or_else(|| vec![None])
566 } else {
567 vec![None]
568 };
569 for variant in &plan.variants {
570 if ctx
571 .variant_id
572 .as_ref()
573 .is_some_and(|requested| requested != &variant.variant_id)
574 {
575 continue;
576 }
577 for fold_id in &fold_ids {
578 let seed_root = variant.seed.or(ctx.root_seed);
579 results.extend(self.execute_phase_scope(
580 plan,
581 controllers,
582 ctx,
583 PhaseScope {
584 phase,
585 variant_id: Some(variant.variant_id.clone()),
586 variant: Some(VariantExecutionSpec::from_plan(variant)),
587 fold_id: fold_id.clone(),
588 seed_root,
589 },
590 PhaseScopeResources::default(),
591 )?);
592 }
593 }
594 Ok(results)
595 }
596
597 pub fn execute_campaign_phase_with_data_provider(
598 &self,
599 plan: &ExecutionPlan,
600 controllers: &RuntimeControllerRegistry,
601 data_provider: &dyn RuntimeDataProvider,
602 ctx: &mut RunContext,
603 phase: Phase,
604 ) -> Result<Vec<NodeResult>> {
605 plan.validate()?;
606 let mut results = Vec::new();
607 let fold_ids = if phase == Phase::FitCv {
608 plan.fold_set
609 .as_ref()
610 .map(|fold_set| {
611 fold_set
612 .folds
613 .iter()
614 .map(|fold| Some(fold.fold_id.clone()))
615 .collect::<Vec<_>>()
616 })
617 .unwrap_or_else(|| vec![None])
618 } else {
619 vec![None]
620 };
621 for variant in &plan.variants {
622 if ctx
623 .variant_id
624 .as_ref()
625 .is_some_and(|requested| requested != &variant.variant_id)
626 {
627 continue;
628 }
629 for fold_id in &fold_ids {
630 let seed_root = variant.seed.or(ctx.root_seed);
631 results.extend(self.execute_phase_scope(
632 plan,
633 controllers,
634 ctx,
635 PhaseScope {
636 phase,
637 variant_id: Some(variant.variant_id.clone()),
638 variant: Some(VariantExecutionSpec::from_plan(variant)),
639 fold_id: fold_id.clone(),
640 seed_root,
641 },
642 PhaseScopeResources {
643 data_provider: Some(data_provider),
644 ..Default::default()
645 },
646 )?);
647 }
648 }
649 Ok(results)
650 }
651
652 pub fn execute_campaign_phase_with_data_provider_and_artifact_store(
653 &self,
654 plan: &ExecutionPlan,
655 controllers: &RuntimeControllerRegistry,
656 data_provider: &dyn RuntimeDataProvider,
657 artifact_store: &mut InMemoryArtifactStore,
658 ctx: &mut RunContext,
659 phase: Phase,
660 ) -> Result<Vec<NodeResult>> {
661 plan.validate()?;
662 let mut results = Vec::new();
663 let fold_ids = if phase == Phase::FitCv {
664 plan.fold_set
665 .as_ref()
666 .map(|fold_set| {
667 fold_set
668 .folds
669 .iter()
670 .map(|fold| Some(fold.fold_id.clone()))
671 .collect::<Vec<_>>()
672 })
673 .unwrap_or_else(|| vec![None])
674 } else {
675 vec![None]
676 };
677 for variant in &plan.variants {
678 if ctx
679 .variant_id
680 .as_ref()
681 .is_some_and(|requested| requested != &variant.variant_id)
682 {
683 continue;
684 }
685 for fold_id in &fold_ids {
686 let seed_root = variant.seed.or(ctx.root_seed);
687 results.extend(self.execute_phase_scope(
688 plan,
689 controllers,
690 ctx,
691 PhaseScope {
692 phase,
693 variant_id: Some(variant.variant_id.clone()),
694 variant: Some(VariantExecutionSpec::from_plan(variant)),
695 fold_id: fold_id.clone(),
696 seed_root,
697 },
698 PhaseScopeResources {
699 data_provider: Some(data_provider),
700 artifact_store: Some(&mut *artifact_store),
701 ..Default::default()
702 },
703 )?);
704 }
705 }
706 Ok(results)
707 }
708
709 pub fn execute_bundle_replay(
710 &self,
711 replay: BundleReplayExecution<'_>,
712 ctx: &mut RunContext,
713 ) -> Result<Vec<NodeResult>> {
714 replay.bundle.validate_against_plan(replay.plan)?;
715 replay
716 .replay_request
717 .validate_for_bundle_with_prediction_cache_store(
718 replay.bundle,
719 replay.prediction_cache_store.is_some(),
720 )?;
721 replay
722 .bundle
723 .validate_replay_envelopes(replay.data_envelopes)?;
724 let prediction_cache_contracts = if replay.replay_request.phase == Phase::Refit {
725 Some(replay_prediction_cache_contracts(replay.bundle)?)
726 } else {
727 None
728 };
729 if replay.replay_request.phase == Phase::Refit {
730 preload_replay_prediction_cache_store(
731 replay.bundle,
732 replay.prediction_cache_store,
733 ctx,
734 )?;
735 }
736 let replay_artifacts = materialize_replay_artifact_handles(
737 replay.plan,
738 replay.bundle,
739 replay.replay_request,
740 replay.artifact_store,
741 ctx,
742 )?;
743 let selected_variant = replay
744 .bundle
745 .selected_variant_id
746 .as_ref()
747 .map(|selected| {
748 replay
749 .plan
750 .variants
751 .iter()
752 .find(|variant| &variant.variant_id == selected)
753 .map(VariantExecutionSpec::from_plan)
754 .ok_or_else(|| {
755 DagMlError::RuntimeValidation(format!(
756 "bundle `{}` selected unknown variant `{selected}`",
757 replay.bundle.bundle_id
758 ))
759 })
760 })
761 .transpose()?;
762 let seed_root = selected_variant
763 .as_ref()
764 .and_then(|variant| variant.seed)
765 .or(ctx.root_seed);
766
767 self.execute_phase_scope(
768 replay.plan,
769 replay.controllers,
770 ctx,
771 PhaseScope {
772 phase: replay.replay_request.phase,
773 variant_id: replay.bundle.selected_variant_id.clone(),
774 variant: selected_variant,
775 fold_id: None,
776 seed_root,
777 },
778 PhaseScopeResources {
779 data_provider: Some(replay.data_provider),
780 replay_artifact_handles: Some(&replay_artifacts.handles),
781 replay_artifact_inputs: Some(&replay_artifacts.inputs),
782 replay_bundle_id: Some(&replay.bundle.bundle_id),
783 data_envelopes: Some(replay.data_envelopes),
784 prediction_cache_store: replay.prediction_cache_store,
785 prediction_cache_contracts: prediction_cache_contracts.as_ref(),
786 ..Default::default()
787 },
788 )
789 }
790
791 fn execute_phase_scope(
792 &self,
793 plan: &ExecutionPlan,
794 controllers: &RuntimeControllerRegistry,
795 ctx: &mut RunContext,
796 scope: PhaseScope,
797 mut resources: PhaseScopeResources<'_>,
798 ) -> Result<Vec<NodeResult>> {
799 let _phase_span = crate::observability::phase_span(
800 ctx.run_id.as_str(),
801 plan.id.as_str(),
802 scope.phase.as_str(),
803 scope.variant_id.as_ref().map(VariantId::as_str),
804 scope.fold_id.as_ref().map(FoldId::as_str),
805 )
806 .entered();
807 let mut results = Vec::new();
808 let mut output_handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
809 let mut output_data_views =
810 BTreeMap::<NodeId, BTreeMap<String, DataProviderViewSpec>>::new();
811 let mut input_lineage = BTreeMap::<NodeId, LineageId>::new();
812
813 for level in plan.node_parallel_levels_for_phase(scope.phase)? {
814 for node_id in &level {
815 let node_plan = plan
816 .node_plans
817 .get(node_id)
818 .expect("execution plan was validated");
819 if let Some(reduction) = merge_reduction_mode(plan, node_plan) {
826 if let Some(mut result) =
827 reassemble_branch_merge(plan, node_plan, ctx, &scope, reduction)?
828 {
829 let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
830 let task = NodeTask {
831 inner_fold_set: None,
832 run_id: ctx.run_id.clone(),
833 node_plan: task_node_plan.clone(),
834 phase: scope.phase,
835 variant_id: scope.variant_id.clone(),
836 variant: scope.variant.clone(),
837 fold_id: scope.fold_id.clone(),
838 branch_path: Vec::new(),
839 input_handles: BTreeMap::new(),
840 data_views: BTreeMap::new(),
841 prediction_inputs: BTreeMap::new(),
842 artifact_inputs: BTreeMap::new(),
843 required_loss_attestations: NodeTask::required_loss_attestations_for(
844 &task_node_plan,
845 scope.phase,
846 )?,
847 fit_influence: FitInfluenceTask::default(),
848 seed: None,
849 };
850 normalize_result_prediction_ports(plan, &task, &mut result)?;
851 result.validate_for_task(&task)?;
852 for prediction in &result.predictions {
853 ctx.prediction_store.append(prediction.clone())?;
854 }
855 apply_result_scoring(
856 &result,
857 &mut ctx.score_collector,
858 &mut ctx.regression_target_records,
859 )?;
860 ctx.lineage.record(result.lineage.clone())?;
861 output_handles.insert(node_id.clone(), result.outputs.clone());
862 input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
863 results.push(result);
864 }
865 continue;
866 }
867 let controller = controllers.get(&node_plan.controller_id).ok_or_else(|| {
868 DagMlError::RuntimeValidation(format!(
869 "runtime controller `{}` is not registered",
870 node_plan.controller_id
871 ))
872 })?;
873 let collected_inputs = collect_input_handles(
874 plan,
875 node_plan,
876 &output_handles,
877 &output_data_views,
878 &resources,
879 ctx,
880 &scope,
881 )?;
882 if collected_inputs.skip_node {
883 continue;
884 }
885 let mut input_handles = collected_inputs.handles;
886 let mut artifact_inputs = BTreeMap::new();
887 if let Some(node_artifact_handles) = resources
888 .replay_artifact_handles
889 .and_then(|handles| handles.get(node_id))
890 {
891 for (key, handle) in node_artifact_handles {
892 if input_handles.insert(key.clone(), handle.clone()).is_some() {
893 return Err(DagMlError::RuntimeValidation(format!(
894 "node `{node_id}` received duplicate replay artifact input `{key}`"
895 )));
896 }
897 }
898 }
899 if let Some(node_artifact_inputs) = resources
900 .replay_artifact_inputs
901 .and_then(|inputs| inputs.get(node_id))
902 {
903 for (key, spec) in node_artifact_inputs {
904 if artifact_inputs.insert(key.clone(), spec.clone()).is_some() {
905 return Err(DagMlError::RuntimeValidation(format!(
906 "node `{node_id}` received duplicate replay artifact metadata `{key}`"
907 )));
908 }
909 }
910 }
911 let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
912 let inner_fold_set = inner_fold_set_for_scope(
913 &plan.campaign,
914 plan.fold_set.as_ref(),
915 node_plan,
916 &scope,
917 )?;
918 let fit_influence = fit_influence_task_for_node(
919 plan,
920 &task_node_plan,
921 &collected_inputs.data_views,
922 )?;
923 let task = NodeTask {
924 inner_fold_set,
925 run_id: ctx.run_id.clone(),
926 node_plan: task_node_plan.clone(),
927 phase: scope.phase,
928 variant_id: scope.variant_id.clone(),
929 variant: scope.variant.clone(),
930 fold_id: scope.fold_id.clone(),
931 branch_path: Vec::new(),
932 input_handles,
933 data_views: collected_inputs.data_views,
934 prediction_inputs: collected_inputs.prediction_inputs,
935 artifact_inputs,
936 required_loss_attestations: NodeTask::required_loss_attestations_for(
937 &task_node_plan,
938 scope.phase,
939 )?,
940 fit_influence,
941 seed: derive_task_seed(
942 scope.seed_root,
943 scope.variant_id.as_ref(),
944 scope.fold_id.as_ref(),
945 &task_node_plan,
946 scope.phase,
947 ),
948 };
949 let _node_span = crate::observability::node_span(
950 task.run_id.as_str(),
951 plan.id.as_str(),
952 task.phase.as_str(),
953 task.node_plan.node_id.as_str(),
954 task.node_plan.controller_id.as_str(),
955 )
956 .entered();
957 let mut result = if task.node_plan.kind == NodeKind::Tuner {
958 return Err(DagMlError::RuntimeValidation(format!(
959 "tuner node `{}` requires execute_hpo_campaign with an explicit RuntimeHpoExecutionContext",
960 task.node_plan.node_id
961 )));
962 } else {
963 match resources.data_provider {
964 Some(data_provider) => {
965 controller.invoke_with_data_provider(&task, data_provider)?
966 }
967 None => controller.invoke(&task)?,
968 }
969 };
970 record_fit_influence_diagnostic(&task, &mut result);
971 normalize_result_prediction_ports(plan, &task, &mut result)?;
972 result.validate_for_task(&task)?;
973 apply_result_prediction_aggregation(
974 plan,
975 controllers,
976 &task,
977 &mut result,
978 &resources,
979 )?;
980 attach_coordinator_input_lineage(
981 &mut result,
982 plan,
983 &task.node_plan.node_id,
984 &input_lineage,
985 )?;
986 if let Some(store) = resources.artifact_store.as_deref_mut() {
987 if scope.phase == Phase::Refit {
988 store.capture_refit_artifacts(&task, &result)?;
989 }
990 }
991 for prediction in &result.predictions {
992 ctx.prediction_store.append(prediction.clone())?;
993 }
994 for prediction in &result.aggregated_predictions {
995 ctx.aggregated_prediction_store.append(prediction.clone())?;
996 }
997 apply_result_scoring(
998 &result,
999 &mut ctx.score_collector,
1000 &mut ctx.regression_target_records,
1001 )?;
1002 ctx.lineage.record(result.lineage.clone())?;
1003 let data_views = derive_output_data_views(plan, &task, &result)?;
1004 output_handles.insert(node_id.clone(), result.outputs.clone());
1005 output_data_views.insert(node_id.clone(), data_views);
1006 input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
1007 results.push(result);
1008 }
1009 }
1010
1011 Ok(results)
1012 }
1013}
1014
1015impl ParallelScheduler {
1016 pub fn execute_phase(
1017 &self,
1018 plan: &ExecutionPlan,
1019 controllers: &RuntimeControllerRegistry,
1020 ctx: &mut RunContext,
1021 phase: Phase,
1022 ) -> Result<Vec<NodeResult>> {
1023 plan.validate()?;
1024 let variant_id = ctx.variant_id.clone();
1025 let seed_root = ctx.root_seed;
1026 self.execute_phase_scope(
1027 plan,
1028 controllers,
1029 ctx,
1030 PhaseScope {
1031 phase,
1032 variant_id,
1033 variant: None,
1034 fold_id: None,
1035 seed_root,
1036 },
1037 PhaseScopeResources::default(),
1038 )
1039 }
1040
1041 pub fn execute_phase_with_data_provider(
1042 &self,
1043 plan: &ExecutionPlan,
1044 controllers: &RuntimeControllerRegistry,
1045 data_provider: &dyn RuntimeDataProvider,
1046 ctx: &mut RunContext,
1047 phase: Phase,
1048 ) -> Result<Vec<NodeResult>> {
1049 plan.validate()?;
1050 let variant_id = ctx.variant_id.clone();
1051 let seed_root = ctx.root_seed;
1052 self.execute_phase_scope(
1053 plan,
1054 controllers,
1055 ctx,
1056 PhaseScope {
1057 phase,
1058 variant_id,
1059 variant: None,
1060 fold_id: None,
1061 seed_root,
1062 },
1063 PhaseScopeResources {
1064 data_provider: Some(data_provider),
1065 ..Default::default()
1066 },
1067 )
1068 }
1069
1070 pub fn execute_campaign_phase(
1071 &self,
1072 plan: &ExecutionPlan,
1073 controllers: &RuntimeControllerRegistry,
1074 ctx: &mut RunContext,
1075 phase: Phase,
1076 ) -> Result<Vec<NodeResult>> {
1077 plan.validate()?;
1078 let mut results = Vec::new();
1079 let fold_ids = if phase == Phase::FitCv {
1080 plan.fold_set
1081 .as_ref()
1082 .map(|fold_set| {
1083 fold_set
1084 .folds
1085 .iter()
1086 .map(|fold| Some(fold.fold_id.clone()))
1087 .collect::<Vec<_>>()
1088 })
1089 .unwrap_or_else(|| vec![None])
1090 } else {
1091 vec![None]
1092 };
1093 for variant in &plan.variants {
1094 if ctx
1095 .variant_id
1096 .as_ref()
1097 .is_some_and(|requested| requested != &variant.variant_id)
1098 {
1099 continue;
1100 }
1101 for fold_id in &fold_ids {
1102 let seed_root = variant.seed.or(ctx.root_seed);
1103 results.extend(self.execute_phase_scope(
1104 plan,
1105 controllers,
1106 ctx,
1107 PhaseScope {
1108 phase,
1109 variant_id: Some(variant.variant_id.clone()),
1110 variant: Some(VariantExecutionSpec::from_plan(variant)),
1111 fold_id: fold_id.clone(),
1112 seed_root,
1113 },
1114 PhaseScopeResources::default(),
1115 )?);
1116 }
1117 }
1118 Ok(results)
1119 }
1120
1121 pub fn execute_campaign_phase_with_data_provider(
1122 &self,
1123 plan: &ExecutionPlan,
1124 controllers: &RuntimeControllerRegistry,
1125 data_provider: &dyn RuntimeDataProvider,
1126 ctx: &mut RunContext,
1127 phase: Phase,
1128 ) -> Result<Vec<NodeResult>> {
1129 plan.validate()?;
1130 let mut results = Vec::new();
1131 let fold_ids = if phase == Phase::FitCv {
1132 plan.fold_set
1133 .as_ref()
1134 .map(|fold_set| {
1135 fold_set
1136 .folds
1137 .iter()
1138 .map(|fold| Some(fold.fold_id.clone()))
1139 .collect::<Vec<_>>()
1140 })
1141 .unwrap_or_else(|| vec![None])
1142 } else {
1143 vec![None]
1144 };
1145 for variant in &plan.variants {
1146 if ctx
1147 .variant_id
1148 .as_ref()
1149 .is_some_and(|requested| requested != &variant.variant_id)
1150 {
1151 continue;
1152 }
1153 for fold_id in &fold_ids {
1154 let seed_root = variant.seed.or(ctx.root_seed);
1155 results.extend(self.execute_phase_scope(
1156 plan,
1157 controllers,
1158 ctx,
1159 PhaseScope {
1160 phase,
1161 variant_id: Some(variant.variant_id.clone()),
1162 variant: Some(VariantExecutionSpec::from_plan(variant)),
1163 fold_id: fold_id.clone(),
1164 seed_root,
1165 },
1166 PhaseScopeResources {
1167 data_provider: Some(data_provider),
1168 ..Default::default()
1169 },
1170 )?);
1171 }
1172 }
1173 Ok(results)
1174 }
1175
1176 pub fn execute_campaign_phase_with_data_provider_and_artifact_store(
1177 &self,
1178 plan: &ExecutionPlan,
1179 controllers: &RuntimeControllerRegistry,
1180 data_provider: &dyn RuntimeDataProvider,
1181 artifact_store: &mut InMemoryArtifactStore,
1182 ctx: &mut RunContext,
1183 phase: Phase,
1184 ) -> Result<Vec<NodeResult>> {
1185 plan.validate()?;
1186 let mut results = Vec::new();
1187 let fold_ids = if phase == Phase::FitCv {
1188 plan.fold_set
1189 .as_ref()
1190 .map(|fold_set| {
1191 fold_set
1192 .folds
1193 .iter()
1194 .map(|fold| Some(fold.fold_id.clone()))
1195 .collect::<Vec<_>>()
1196 })
1197 .unwrap_or_else(|| vec![None])
1198 } else {
1199 vec![None]
1200 };
1201 for variant in &plan.variants {
1202 if ctx
1203 .variant_id
1204 .as_ref()
1205 .is_some_and(|requested| requested != &variant.variant_id)
1206 {
1207 continue;
1208 }
1209 for fold_id in &fold_ids {
1210 let seed_root = variant.seed.or(ctx.root_seed);
1211 results.extend(self.execute_phase_scope(
1212 plan,
1213 controllers,
1214 ctx,
1215 PhaseScope {
1216 phase,
1217 variant_id: Some(variant.variant_id.clone()),
1218 variant: Some(VariantExecutionSpec::from_plan(variant)),
1219 fold_id: fold_id.clone(),
1220 seed_root,
1221 },
1222 PhaseScopeResources {
1223 data_provider: Some(data_provider),
1224 artifact_store: Some(&mut *artifact_store),
1225 ..Default::default()
1226 },
1227 )?);
1228 }
1229 }
1230 Ok(results)
1231 }
1232
1233 pub fn execute_bundle_replay(
1234 &self,
1235 replay: BundleReplayExecution<'_>,
1236 ctx: &mut RunContext,
1237 ) -> Result<Vec<NodeResult>> {
1238 replay.bundle.validate_against_plan(replay.plan)?;
1239 replay
1240 .replay_request
1241 .validate_for_bundle_with_prediction_cache_store(
1242 replay.bundle,
1243 replay.prediction_cache_store.is_some(),
1244 )?;
1245 replay
1246 .bundle
1247 .validate_replay_envelopes(replay.data_envelopes)?;
1248 let prediction_cache_contracts = if replay.replay_request.phase == Phase::Refit {
1249 Some(replay_prediction_cache_contracts(replay.bundle)?)
1250 } else {
1251 None
1252 };
1253 if replay.replay_request.phase == Phase::Refit {
1254 preload_replay_prediction_cache_store(
1255 replay.bundle,
1256 replay.prediction_cache_store,
1257 ctx,
1258 )?;
1259 }
1260 let replay_artifacts = materialize_replay_artifact_handles(
1261 replay.plan,
1262 replay.bundle,
1263 replay.replay_request,
1264 replay.artifact_store,
1265 ctx,
1266 )?;
1267 let selected_variant = replay
1268 .bundle
1269 .selected_variant_id
1270 .as_ref()
1271 .map(|selected| {
1272 replay
1273 .plan
1274 .variants
1275 .iter()
1276 .find(|variant| &variant.variant_id == selected)
1277 .map(VariantExecutionSpec::from_plan)
1278 .ok_or_else(|| {
1279 DagMlError::RuntimeValidation(format!(
1280 "bundle `{}` selected unknown variant `{selected}`",
1281 replay.bundle.bundle_id
1282 ))
1283 })
1284 })
1285 .transpose()?;
1286 let seed_root = selected_variant
1287 .as_ref()
1288 .and_then(|variant| variant.seed)
1289 .or(ctx.root_seed);
1290
1291 self.execute_phase_scope(
1292 replay.plan,
1293 replay.controllers,
1294 ctx,
1295 PhaseScope {
1296 phase: replay.replay_request.phase,
1297 variant_id: replay.bundle.selected_variant_id.clone(),
1298 variant: selected_variant,
1299 fold_id: None,
1300 seed_root,
1301 },
1302 PhaseScopeResources {
1303 data_provider: Some(replay.data_provider),
1304 replay_artifact_handles: Some(&replay_artifacts.handles),
1305 replay_artifact_inputs: Some(&replay_artifacts.inputs),
1306 replay_bundle_id: Some(&replay.bundle.bundle_id),
1307 data_envelopes: Some(replay.data_envelopes),
1308 prediction_cache_store: replay.prediction_cache_store,
1309 prediction_cache_contracts: prediction_cache_contracts.as_ref(),
1310 ..Default::default()
1311 },
1312 )
1313 }
1314
1315 fn execute_phase_scope(
1316 &self,
1317 plan: &ExecutionPlan,
1318 controllers: &RuntimeControllerRegistry,
1319 ctx: &mut RunContext,
1320 scope: PhaseScope,
1321 mut resources: PhaseScopeResources<'_>,
1322 ) -> Result<Vec<NodeResult>> {
1323 let phase_span = crate::observability::phase_span(
1327 ctx.run_id.as_str(),
1328 plan.id.as_str(),
1329 scope.phase.as_str(),
1330 scope.variant_id.as_ref().map(VariantId::as_str),
1331 scope.fold_id.as_ref().map(FoldId::as_str),
1332 );
1333 let _phase_entered = phase_span.clone().entered();
1334 let plan_id = plan.id.as_str();
1336 plan.validate_parallel_controller_capabilities(self.max_workers, scope.phase)?;
1337 let mut results = Vec::new();
1338 let mut output_handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
1339 let mut output_data_views =
1340 BTreeMap::<NodeId, BTreeMap<String, DataProviderViewSpec>>::new();
1341 let mut input_lineage = BTreeMap::<NodeId, LineageId>::new();
1342
1343 for level in plan.node_parallel_levels_for_phase(scope.phase)? {
1344 let mut prepared = Vec::<PreparedNodeTask>::new();
1345 let mut merge_nodes = Vec::<(NodeId, MergeReduction)>::new();
1352 for node_id in &level {
1353 let node_plan = plan
1354 .node_plans
1355 .get(node_id)
1356 .expect("execution plan was validated");
1357 if let Some(reduction) = merge_reduction_mode(plan, node_plan) {
1358 merge_nodes.push((node_id.clone(), reduction));
1359 continue;
1360 }
1361 let collected_inputs = collect_input_handles(
1362 plan,
1363 node_plan,
1364 &output_handles,
1365 &output_data_views,
1366 &resources,
1367 ctx,
1368 &scope,
1369 )?;
1370 if collected_inputs.skip_node {
1371 continue;
1372 }
1373 let mut input_handles = collected_inputs.handles;
1374 let mut artifact_inputs = BTreeMap::new();
1375 if let Some(node_artifact_handles) = resources
1376 .replay_artifact_handles
1377 .and_then(|handles| handles.get(node_id))
1378 {
1379 for (key, handle) in node_artifact_handles {
1380 if input_handles.insert(key.clone(), handle.clone()).is_some() {
1381 return Err(DagMlError::RuntimeValidation(format!(
1382 "node `{node_id}` received duplicate replay artifact input `{key}`"
1383 )));
1384 }
1385 }
1386 }
1387 if let Some(node_artifact_inputs) = resources
1388 .replay_artifact_inputs
1389 .and_then(|inputs| inputs.get(node_id))
1390 {
1391 for (key, spec) in node_artifact_inputs {
1392 if artifact_inputs.insert(key.clone(), spec.clone()).is_some() {
1393 return Err(DagMlError::RuntimeValidation(format!(
1394 "node `{node_id}` received duplicate replay artifact metadata `{key}`"
1395 )));
1396 }
1397 }
1398 }
1399 let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
1400 let inner_fold_set = inner_fold_set_for_scope(
1401 &plan.campaign,
1402 plan.fold_set.as_ref(),
1403 node_plan,
1404 &scope,
1405 )?;
1406 let fit_influence = fit_influence_task_for_node(
1407 plan,
1408 &task_node_plan,
1409 &collected_inputs.data_views,
1410 )?;
1411 prepared.push(PreparedNodeTask {
1412 node_id: node_id.clone(),
1413 task: NodeTask {
1414 inner_fold_set,
1415 run_id: ctx.run_id.clone(),
1416 node_plan: task_node_plan.clone(),
1417 phase: scope.phase,
1418 variant_id: scope.variant_id.clone(),
1419 variant: scope.variant.clone(),
1420 fold_id: scope.fold_id.clone(),
1421 branch_path: Vec::new(),
1422 input_handles,
1423 data_views: collected_inputs.data_views,
1424 prediction_inputs: collected_inputs.prediction_inputs,
1425 artifact_inputs,
1426 required_loss_attestations: NodeTask::required_loss_attestations_for(
1427 &task_node_plan,
1428 scope.phase,
1429 )?,
1430 fit_influence,
1431 seed: derive_task_seed(
1432 scope.seed_root,
1433 scope.variant_id.as_ref(),
1434 scope.fold_id.as_ref(),
1435 &task_node_plan,
1436 scope.phase,
1437 ),
1438 },
1439 });
1440 }
1441
1442 for chunk in prepared.chunks(self.max_workers) {
1443 let chunk_results = std::thread::scope(
1444 |thread_scope| -> Result<Vec<NodeResult>> {
1445 let mut handles = Vec::with_capacity(chunk.len());
1446 for prepared_task in chunk {
1447 let controller = controllers
1448 .get(&prepared_task.task.node_plan.controller_id)
1449 .ok_or_else(|| {
1450 DagMlError::RuntimeValidation(format!(
1451 "runtime controller `{}` is not registered",
1452 prepared_task.task.node_plan.controller_id
1453 ))
1454 })?;
1455 let worker_span = phase_span.clone();
1456 handles.push(thread_scope.spawn(move || {
1457 let _worker_span = worker_span.entered();
1458 let _node_span = crate::observability::node_span(
1459 prepared_task.task.run_id.as_str(),
1460 plan_id,
1461 prepared_task.task.phase.as_str(),
1462 prepared_task.task.node_plan.node_id.as_str(),
1463 prepared_task.task.node_plan.controller_id.as_str(),
1464 )
1465 .entered();
1466 let mut result =
1467 if prepared_task.task.node_plan.kind == NodeKind::Tuner {
1468 return Err(DagMlError::RuntimeValidation(format!(
1469 "tuner node `{}` requires execute_hpo_campaign with an explicit RuntimeHpoExecutionContext",
1470 prepared_task.task.node_plan.node_id
1471 )));
1472 } else {
1473 controller.invoke(&prepared_task.task)?
1479 };
1480 record_fit_influence_diagnostic(&prepared_task.task, &mut result);
1481 normalize_result_prediction_ports(
1482 plan,
1483 &prepared_task.task,
1484 &mut result,
1485 )?;
1486 result.validate_for_task(&prepared_task.task)?;
1487 Ok(result)
1488 }));
1489 }
1490 handles
1491 .into_iter()
1492 .map(|handle| {
1493 handle.join().map_err(|_| {
1494 DagMlError::RuntimeValidation(
1495 "parallel scheduler worker panicked".to_string(),
1496 )
1497 })?
1498 })
1499 .collect()
1500 },
1501 )?;
1502
1503 for (prepared_task, mut result) in chunk.iter().zip(chunk_results) {
1504 apply_result_prediction_aggregation(
1505 plan,
1506 controllers,
1507 &prepared_task.task,
1508 &mut result,
1509 &resources,
1510 )?;
1511 attach_coordinator_input_lineage(
1512 &mut result,
1513 plan,
1514 &prepared_task.task.node_plan.node_id,
1515 &input_lineage,
1516 )?;
1517 if let Some(store) = resources.artifact_store.as_deref_mut() {
1518 if scope.phase == Phase::Refit {
1519 store.capture_refit_artifacts(&prepared_task.task, &result)?;
1520 }
1521 }
1522 for prediction in &result.predictions {
1523 ctx.prediction_store.append(prediction.clone())?;
1524 }
1525 for prediction in &result.aggregated_predictions {
1526 ctx.aggregated_prediction_store.append(prediction.clone())?;
1527 }
1528 apply_result_scoring(
1529 &result,
1530 &mut ctx.score_collector,
1531 &mut ctx.regression_target_records,
1532 )?;
1533 ctx.lineage.record(result.lineage.clone())?;
1534 let data_views = derive_output_data_views(plan, &prepared_task.task, &result)?;
1535 output_handles.insert(prepared_task.node_id.clone(), result.outputs.clone());
1536 output_data_views.insert(prepared_task.node_id.clone(), data_views);
1537 input_lineage.insert(
1538 prepared_task.node_id.clone(),
1539 result.lineage.record_id.clone(),
1540 );
1541 results.push(result);
1542 }
1543 }
1544
1545 for (node_id, reduction) in &merge_nodes {
1550 let node_plan = plan
1551 .node_plans
1552 .get(node_id)
1553 .expect("execution plan was validated");
1554 if let Some(mut result) =
1555 reassemble_branch_merge(plan, node_plan, ctx, &scope, *reduction)?
1556 {
1557 let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
1558 let task = NodeTask {
1559 inner_fold_set: None,
1560 run_id: ctx.run_id.clone(),
1561 node_plan: task_node_plan.clone(),
1562 phase: scope.phase,
1563 variant_id: scope.variant_id.clone(),
1564 variant: scope.variant.clone(),
1565 fold_id: scope.fold_id.clone(),
1566 branch_path: Vec::new(),
1567 input_handles: BTreeMap::new(),
1568 data_views: BTreeMap::new(),
1569 prediction_inputs: BTreeMap::new(),
1570 artifact_inputs: BTreeMap::new(),
1571 required_loss_attestations: NodeTask::required_loss_attestations_for(
1572 &task_node_plan,
1573 scope.phase,
1574 )?,
1575 fit_influence: FitInfluenceTask::default(),
1576 seed: None,
1577 };
1578 normalize_result_prediction_ports(plan, &task, &mut result)?;
1579 result.validate_for_task(&task)?;
1580 for prediction in &result.predictions {
1581 ctx.prediction_store.append(prediction.clone())?;
1582 }
1583 apply_result_scoring(
1584 &result,
1585 &mut ctx.score_collector,
1586 &mut ctx.regression_target_records,
1587 )?;
1588 ctx.lineage.record(result.lineage.clone())?;
1589 output_handles.insert(node_id.clone(), result.outputs.clone());
1590 input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
1591 results.push(result);
1592 }
1593 }
1594 }
1595
1596 Ok(results)
1597 }
1598}
1599
1600#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1601enum HpoTrialTerminalState {
1602 Completed,
1603 Pruned,
1604 Failed,
1605}
1606
1607fn validate_hpo_checkpoint_result(
1608 checkpoint: &RuntimeHpoCheckpointResult,
1609 hpo: &RuntimeHpoExecutionContext,
1610 trial_variants: &BTreeMap<i64, VariantId>,
1611 terminal_trials: &BTreeMap<i64, HpoTrialTerminalState>,
1612 history_at_start: u32,
1613) -> Result<()> {
1614 checkpoint.artifact.validate().map_err(|error| {
1615 DagMlError::RuntimeValidation(format!(
1616 "runtime HPO checkpoint artifact is invalid: {error}"
1617 ))
1618 })?;
1619 if checkpoint.operation_id != hpo.operation_id
1620 || checkpoint.controller_id != hpo.controller_id
1621 || checkpoint.target_node_id != hpo.target_node_id
1622 || checkpoint.provenance != hpo.provenance
1623 {
1624 return Err(DagMlError::RuntimeValidation(
1625 "runtime HPO checkpoint provenance does not exactly match its execution context"
1626 .to_string(),
1627 ));
1628 }
1629 let proposed_count = u32::try_from(trial_variants.len()).map_err(|_| {
1630 DagMlError::RuntimeValidation(
1631 "runtime HPO scheduler proposal count does not fit u32".to_string(),
1632 )
1633 })?;
1634 if checkpoint.trial_history_len != hpo.trial_budget_total
1635 || checkpoint.trial_history_len < history_at_start
1636 || checkpoint.trial_history_len - history_at_start != proposed_count
1637 {
1638 return Err(DagMlError::RuntimeValidation(
1639 "runtime HPO checkpoint native history is inconsistent with scheduler-observed trials"
1640 .to_string(),
1641 ));
1642 }
1643 if checkpoint.artifact.binding.controller_id != hpo.controller_id.as_str()
1644 || checkpoint.artifact.binding.controller_id != hpo.study.controller_id
1645 || checkpoint.artifact.binding.study_id != hpo.study.study_id
1646 || checkpoint.artifact.methods_abi != hpo.study.methods_abi
1647 {
1648 return Err(DagMlError::RuntimeValidation(
1649 "runtime HPO checkpoint binding/controller/study does not match the active tuner"
1650 .to_string(),
1651 ));
1652 }
1653 let expected_search_space = hpo.study.search_space.fingerprint().map_err(|error| {
1654 DagMlError::RuntimeValidation(format!(
1655 "runtime HPO cannot fingerprint the configured search space: {error}"
1656 ))
1657 })?;
1658 if checkpoint.artifact.binding.search_space_fingerprint != expected_search_space {
1659 return Err(DagMlError::RuntimeValidation(
1660 "runtime HPO checkpoint search-space binding does not match the active study"
1661 .to_string(),
1662 ));
1663 }
1664
1665 let completed_trial_ids = terminal_trials
1666 .iter()
1667 .filter_map(|(trial_id, state)| {
1668 (*state == HpoTrialTerminalState::Completed).then_some(*trial_id)
1669 })
1670 .collect::<BTreeSet<_>>();
1671 let mut proposal_trial_ids = BTreeSet::new();
1672 for proposal in &checkpoint.completed_proposals {
1673 if !proposal_trial_ids.insert(proposal.trial_id) {
1674 return Err(DagMlError::RuntimeValidation(format!(
1675 "runtime HPO checkpoint has duplicate completed proposal for trial `{}`",
1676 proposal.trial_id
1677 )));
1678 }
1679 if trial_variants.get(&proposal.trial_id) != Some(&proposal.variant.variant_id) {
1680 return Err(DagMlError::RuntimeValidation(format!(
1681 "runtime HPO checkpoint proposal for trial `{}` does not exactly match its scheduler proposal",
1682 proposal.trial_id
1683 )));
1684 }
1685 }
1686 if proposal_trial_ids != completed_trial_ids {
1687 return Err(DagMlError::RuntimeValidation(
1688 "runtime HPO checkpoint proposals must cover exactly the completed trials".to_string(),
1689 ));
1690 }
1691
1692 let mut report_trial_ids = BTreeSet::new();
1693 for completed in &checkpoint.completed_reports {
1694 if !report_trial_ids.insert(completed.trial_id) {
1695 return Err(DagMlError::RuntimeValidation(format!(
1696 "runtime HPO checkpoint has duplicate completed report for trial `{}`",
1697 completed.trial_id
1698 )));
1699 }
1700 if trial_variants.get(&completed.trial_id) != Some(&completed.variant_id)
1701 || !proposal_trial_ids.contains(&completed.trial_id)
1702 {
1703 return Err(DagMlError::RuntimeValidation(format!(
1704 "runtime HPO checkpoint report for trial `{}` does not match a completed proposal",
1705 completed.trial_id
1706 )));
1707 }
1708 let report = &completed.report;
1709 if report.producer_node != hpo.selection.producer_node
1710 || report.producer_port.as_deref() != Some(hpo.selection.producer_port.as_str())
1711 || report.partition != PredictionPartition::Validation
1712 || report
1713 .fold_id
1714 .as_ref()
1715 .is_none_or(|fold| fold.as_str() != "avg")
1716 || report.variant_id.as_ref() != Some(&completed.variant_id)
1717 || !report
1718 .metrics
1719 .get(hpo.selection.metric.name())
1720 .is_some_and(|score| score.is_finite())
1721 {
1722 return Err(DagMlError::RuntimeValidation(format!(
1723 "runtime HPO checkpoint report for trial `{}` is not its one finite target OOF average",
1724 completed.trial_id
1725 )));
1726 }
1727 }
1728 if report_trial_ids != completed_trial_ids {
1729 return Err(DagMlError::RuntimeValidation(
1730 "runtime HPO checkpoint reports must cover exactly one OOF average per completed trial"
1731 .to_string(),
1732 ));
1733 }
1734 Ok(())
1735}
1736
1737pub(crate) struct PreparedNodeTask {
1738 pub(crate) node_id: NodeId,
1739 pub(crate) task: NodeTask,
1740}
1741
1742#[cfg(test)]
1745#[allow(clippy::items_after_test_module)]
1746mod hpo_scheduler_tests {
1747 use std::collections::{BTreeMap, BTreeSet};
1748 use std::sync::{Arc, Mutex};
1749
1750 use sha2::{Digest, Sha256};
1751
1752 use super::*;
1753 use crate::controller::{
1754 ArtifactPolicy, ControllerCapability, ControllerFitScope, ControllerManifest,
1755 ControllerRegistry, RngPolicy,
1756 };
1757 use crate::data::InMemoryDataProvider;
1758 use crate::fold::{FoldAssignment, FoldPartitionMode};
1759 use crate::graph::{GraphInterface, GraphSpec, NodeSpec, PortSchema, PortSpec};
1760 use crate::hpo::{
1761 HpoDirection, HpoMetric, HpoOptimizerConfig, HpoParameter, HpoPruner, HpoSampler,
1762 HpoSearchSpace, HpoStudyBinding, MethodsHpoStudyConfig, N4moptCheckpointArtifact,
1763 N4MOPT_ARTIFACT_KIND, N4MOPT_CHECKPOINT_SCHEMA_VERSION, N4MOPT_FORMAT,
1764 };
1765 use crate::metrics::RegressionTargetBlock;
1766 use crate::oof::PredictionBlock;
1767 use crate::plan::{build_execution_plan, SplitInvocation};
1768
1769 struct HpoTestModel {
1770 id: ControllerId,
1771 trace: Arc<Mutex<Vec<String>>>,
1772 }
1773
1774 impl RuntimeController for HpoTestModel {
1775 fn controller_id(&self) -> &ControllerId {
1776 &self.id
1777 }
1778
1779 fn invoke(&self, task: &NodeTask) -> Result<NodeResult> {
1780 self.trace.lock().unwrap().push("model_cv".to_string());
1781 let sample_id = match task.fold_id.as_ref().map(FoldId::as_str) {
1782 Some("fold:0") => SampleId::new("sample:one").unwrap(),
1783 Some("fold:1") => SampleId::new("sample:two").unwrap(),
1784 other => {
1785 return Err(DagMlError::RuntimeValidation(format!(
1786 "HPO test model received unexpected fold {other:?}"
1787 )));
1788 }
1789 };
1790 Ok(NodeResult {
1791 schema_version: None,
1792 node_id: task.node_plan.node_id.clone(),
1793 outputs: BTreeMap::from([(
1794 "prediction".to_string(),
1795 HandleRef {
1796 handle: 2,
1797 kind: HandleKind::Prediction,
1798 owner_controller: self.id.clone(),
1799 },
1800 )]),
1801 predictions: vec![PredictionBlock {
1802 prediction_id: Some(format!("prediction:{}", task.fold_id.as_ref().unwrap())),
1803 producer_node: task.node_plan.node_id.clone(),
1804 producer_port: None,
1805 partition: PredictionPartition::Validation,
1806 fold_id: task.fold_id.clone(),
1807 sample_ids: vec![sample_id.clone()],
1808 values: vec![vec![1.0]],
1809 target_names: vec!["target".to_string()],
1810 }],
1811 observation_predictions: Vec::new(),
1812 aggregated_predictions: Vec::new(),
1813 explanations: Vec::new(),
1814 shape_deltas: Vec::new(),
1815 artifacts: Vec::new(),
1816 artifact_handles: BTreeMap::new(),
1817 fit_influence_diagnostics: Vec::new(),
1818 regression_targets: vec![RegressionTargetBlock {
1819 level: PredictionLevel::Sample,
1820 unit_ids: vec![PredictionUnitId::Sample(sample_id)],
1821 values: vec![vec![1.0]],
1822 target_names: vec!["target".to_string()],
1823 }],
1824 lineage: LineageRecord {
1825 record_id: LineageId::new(format!(
1826 "lineage:hpo-model:{}",
1827 task.fold_id.as_ref().unwrap()
1828 ))
1829 .unwrap(),
1830 run_id: task.run_id.clone(),
1831 node_id: task.node_plan.node_id.clone(),
1832 phase: task.phase,
1833 controller_id: self.id.clone(),
1834 controller_version: task.node_plan.controller_version.clone(),
1835 variant_id: task.variant_id.clone(),
1836 fold_id: task.fold_id.clone(),
1837 branch_path: Vec::new(),
1838 input_lineage: Vec::new(),
1839 artifact_refs: Vec::new(),
1840 params_fingerprint: task.node_plan.params_fingerprint.clone(),
1841 data_model_shape_fingerprint: None,
1842 aggregation_policy_fingerprint: None,
1843 seed: task.seed,
1844 unsafe_flags: BTreeSet::new(),
1845 metrics: BTreeMap::new(),
1846 loss_attestations: Vec::new(),
1847 early_stopping_records: Vec::new(),
1848 },
1849 })
1850 }
1851 }
1852
1853 struct HpoTestTuner {
1854 id: ControllerId,
1855 trace: Arc<Mutex<Vec<String>>>,
1856 history_len: u32,
1857 proposal_count: u32,
1858 }
1859
1860 struct HpoTestSession {
1861 proposals: Vec<RuntimeHpoProposal>,
1862 trace: Arc<Mutex<Vec<String>>>,
1863 checkpoint: N4moptCheckpointArtifact,
1864 history_len: u32,
1865 completed: Option<(i64, f64)>,
1866 }
1867
1868 impl RuntimeController for HpoTestTuner {
1869 fn controller_id(&self) -> &ControllerId {
1870 &self.id
1871 }
1872
1873 fn invoke(&self, task: &NodeTask) -> Result<NodeResult> {
1874 Err(DagMlError::RuntimeValidation(format!(
1875 "HPO test tuner `{}` was dispatched through generic invoke",
1876 task.node_plan.node_id
1877 )))
1878 }
1879
1880 fn create_tuner_session(
1881 &self,
1882 task: &RuntimeHpoCampaignTask,
1883 context: &RuntimeHpoExecutionContext,
1884 ) -> Result<Box<dyn RuntimeTunerSession>> {
1885 assert_eq!(task.operation_id, context.operation_id);
1886 self.trace
1887 .lock()
1888 .unwrap()
1889 .push("session_factory".to_string());
1890 let payload = vec![7_u8];
1891 let proposals = (0..self.proposal_count)
1892 .map(|offset| {
1893 let trial_id = i64::from(self.history_len + offset + 1);
1894 let mut variant = context.base_variant.clone();
1895 if self.history_len != 0 || self.proposal_count != 1 {
1896 variant.variant_id = VariantId::new(format!("hpo:trial:{trial_id}"))
1897 .map_err(|error| DagMlError::RuntimeValidation(error.to_string()))?;
1898 variant.fingerprint = format!("hpo-test-{trial_id}");
1899 }
1900 Ok(RuntimeHpoProposal { trial_id, variant })
1901 })
1902 .collect::<Result<Vec<_>>>()?;
1903 Ok(Box::new(HpoTestSession {
1904 proposals: proposals.into_iter().rev().collect(),
1905 trace: Arc::clone(&self.trace),
1906 history_len: self.history_len,
1907 completed: None,
1908 checkpoint: N4moptCheckpointArtifact {
1909 schema_version: N4MOPT_CHECKPOINT_SCHEMA_VERSION,
1910 artifact_kind: N4MOPT_ARTIFACT_KIND.to_string(),
1911 format: N4MOPT_FORMAT.to_string(),
1912 binding: HpoStudyBinding {
1913 controller_id: context.study.controller_id.clone(),
1914 study_id: context.study.study_id.clone(),
1915 search_space_fingerprint: context
1916 .study
1917 .search_space
1918 .fingerprint()
1919 .map_err(|error| DagMlError::RuntimeValidation(error.to_string()))?,
1920 optimizer_fingerprint: "optimizer:test".to_string(),
1921 },
1922 methods_abi: context.study.methods_abi.clone(),
1923 payload_sha256: format!("{:x}", Sha256::digest(&payload)),
1924 opaque_payload: payload,
1925 },
1926 }))
1927 }
1928 }
1929
1930 impl RuntimeTunerSession for HpoTestSession {
1931 fn trial_history_len(&self) -> Result<u32> {
1932 Ok(self.history_len)
1933 }
1934
1935 fn ask(&mut self) -> Result<Option<RuntimeHpoProposal>> {
1936 self.trace.lock().unwrap().push("ask".to_string());
1937 let proposal = self.proposals.pop();
1938 if proposal.is_some() {
1939 self.history_len += 1;
1940 }
1941 Ok(proposal)
1942 }
1943
1944 fn report_intermediate(
1945 &mut self,
1946 intermediate: RuntimeHpoIntermediate,
1947 ) -> Result<RuntimeHpoIntermediateOutcome> {
1948 assert_eq!(intermediate.step, 0);
1949 assert!(intermediate.score.is_finite());
1950 self.trace.lock().unwrap().push("intermediate".to_string());
1951 Ok(RuntimeHpoIntermediateOutcome::Continue)
1952 }
1953
1954 fn tell(&mut self, trial_id: i64, terminal: RuntimeHpoTerminal) -> Result<()> {
1955 assert!(trial_id > 0);
1956 assert!(
1957 matches!(terminal, RuntimeHpoTerminal::Completed { score } if score.is_finite())
1958 );
1959 self.trace.lock().unwrap().push("tell".to_string());
1960 if let RuntimeHpoTerminal::Completed { score } = terminal {
1961 self.completed = Some((trial_id, score));
1962 }
1963 Ok(())
1964 }
1965
1966 fn checkpoint(&self) -> Result<N4moptCheckpointArtifact> {
1967 self.trace.lock().unwrap().push("checkpoint".to_string());
1968 Ok(self.checkpoint.clone())
1969 }
1970
1971 fn incumbent(
1972 &self,
1973 variants: &BTreeMap<i64, VariantId>,
1974 ) -> Result<Option<RuntimeHpoIncumbent>> {
1975 let Some((trial_id, score)) = self.completed else {
1976 return Ok(None);
1977 };
1978 Ok(Some(RuntimeHpoIncumbent {
1979 trial_id,
1980 score,
1981 metric: "rmse".to_string(),
1982 direction: HpoDirection::Minimize,
1983 variant_id: variants.get(&trial_id).cloned().unwrap(),
1984 }))
1985 }
1986
1987 fn terminal_trial_snapshots(
1988 &self,
1989 variants: &BTreeMap<i64, VariantId>,
1990 ) -> Result<Vec<RuntimeHpoTerminalSnapshot>> {
1991 let (trial_id, score) = self.completed.ok_or_else(|| {
1992 DagMlError::RuntimeValidation("test HPO session has no completed trial".to_string())
1993 })?;
1994 Ok((1..=i64::from(self.history_len))
1995 .map(|id| {
1996 let completed = id == trial_id;
1997 RuntimeHpoTerminalSnapshot {
1998 trial: crate::hpo::HpoTrial {
1999 id,
2000 ask_sequence: id,
2001 terminal_sequence: Some(id),
2002 parameters: BTreeMap::new(),
2003 parameter_order: Vec::new(),
2004 status: if completed {
2005 crate::hpo::HpoTrialStatus::Completed
2006 } else {
2007 crate::hpo::HpoTrialStatus::Failed
2008 },
2009 score: completed.then_some(score),
2010 rung: 0,
2011 duration: 0.0,
2012 intermediates: Vec::new(),
2013 failure: (!completed).then(|| crate::hpo::HpoFailure {
2014 code: "RESTORED_TEST_FAILURE".to_string(),
2015 message: "synthetic restored terminal".to_string(),
2016 retryable: false,
2017 }),
2018 },
2019 variant_id: variants.get(&id).cloned(),
2020 }
2021 })
2022 .collect())
2023 }
2024 }
2025
2026 fn node(id: &str, kind: NodeKind, outputs: Vec<PortSpec>) -> NodeSpec {
2027 NodeSpec {
2028 id: NodeId::new(id).unwrap(),
2029 kind,
2030 operator: None,
2031 params: BTreeMap::new(),
2032 ports: PortSchema {
2033 inputs: Vec::new(),
2034 outputs,
2035 },
2036 metadata: BTreeMap::new(),
2037 seed_label: None,
2038 }
2039 }
2040
2041 fn manifest(id: &str, kind: NodeKind) -> ControllerManifest {
2042 ControllerManifest {
2043 controller_id: ControllerId::new(id).unwrap(),
2044 controller_version: "test".to_string(),
2045 operator_kind: kind,
2046 priority: 0,
2047 supported_phases: BTreeSet::from([Phase::FitCv]),
2048 input_ports: Vec::new(),
2049 output_ports: Vec::new(),
2050 data_requirements: None,
2051 capabilities: BTreeSet::from([
2052 ControllerCapability::Deterministic,
2053 ControllerCapability::EmitsPredictions,
2054 ]),
2055 operator_selectors: Vec::new(),
2056 fit_scope: ControllerFitScope::FoldTrain,
2057 rng_policy: RngPolicy::UsesCoreSeed,
2058 artifact_policy: ArtifactPolicy::Serializable,
2059 }
2060 }
2061
2062 #[test]
2063 fn hpo_campaign_invokes_registered_session_and_routes_oof_feedback() {
2064 let target = NodeId::new("model:score").unwrap();
2065 let graph = GraphSpec {
2066 id: "graph:hpo.scheduler".to_string(),
2067 interface: GraphInterface::default(),
2068 nodes: vec![node(
2069 "model:score",
2070 NodeKind::Model,
2071 vec![PortSpec {
2072 name: "prediction".to_string(),
2073 kind: PortKind::Prediction,
2074 representation: None,
2075 cardinality: crate::graph::PortCardinality::One,
2076 unit_level: None,
2077 alignment_key: None,
2078 target_level: None,
2079 description: String::new(),
2080 }],
2081 )],
2082 edges: Vec::new(),
2083 search_space_fingerprint: None,
2084 metadata: BTreeMap::new(),
2085 };
2086 let fold_set = FoldSet {
2087 id: "folds:hpo".to_string(),
2088 sample_ids: vec![
2089 SampleId::new("sample:one").unwrap(),
2090 SampleId::new("sample:two").unwrap(),
2091 ],
2092 folds: vec![
2093 FoldAssignment {
2094 fold_id: FoldId::new("fold:0").unwrap(),
2095 train_sample_ids: vec![SampleId::new("sample:two").unwrap()],
2096 validation_sample_ids: vec![SampleId::new("sample:one").unwrap()],
2097 metadata: BTreeMap::new(),
2098 },
2099 FoldAssignment {
2100 fold_id: FoldId::new("fold:1").unwrap(),
2101 train_sample_ids: vec![SampleId::new("sample:one").unwrap()],
2102 validation_sample_ids: vec![SampleId::new("sample:two").unwrap()],
2103 metadata: BTreeMap::new(),
2104 },
2105 ],
2106 sample_groups: BTreeMap::new(),
2107 partition_mode: FoldPartitionMode::Partition,
2108 };
2109 let mut registry = ControllerRegistry::new();
2110 registry
2111 .register(manifest("controller:model", NodeKind::Model))
2112 .unwrap();
2113 let plan = build_execution_plan(
2114 "plan:hpo.scheduler",
2115 graph,
2116 CampaignSpec {
2117 inner_cv: None,
2118 id: "campaign:hpo.scheduler".to_string(),
2119 root_seed: Some(13),
2120 leakage_policy: Default::default(),
2121 aggregation_policy: Default::default(),
2122 split_invocation: Some(SplitInvocation {
2123 id: "split:hpo".to_string(),
2124 controller_id: None,
2125 leakage_policy: Default::default(),
2126 params: BTreeMap::new(),
2127 fold_set: Some(fold_set),
2128 }),
2129 generation: Default::default(),
2130 shape_plans: BTreeMap::new(),
2131 data_bindings: BTreeMap::new(),
2132 branch_view_plans: Vec::new(),
2133 metadata: BTreeMap::new(),
2134 },
2135 ®istry,
2136 )
2137 .unwrap();
2138 let trace = Arc::new(Mutex::new(Vec::new()));
2139 let mut controllers = RuntimeControllerRegistry::new();
2140 controllers
2141 .register(Box::new(HpoTestTuner {
2142 id: ControllerId::new("controller:tuner").unwrap(),
2143 trace: Arc::clone(&trace),
2144 history_len: 0,
2145 proposal_count: 1,
2146 }))
2147 .unwrap();
2148 controllers
2149 .register(Box::new(HpoTestModel {
2150 id: ControllerId::new("controller:model").unwrap(),
2151 trace: Arc::clone(&trace),
2152 }))
2153 .unwrap();
2154 let hpo = RuntimeHpoExecutionContext {
2155 operation_id: "hpo:test".to_string(),
2156 controller_id: ControllerId::new("controller:tuner").unwrap(),
2157 target_node_id: target.clone(),
2158 base_variant: plan.variants[0].clone(),
2159 trial_budget_total: 1,
2160 study: MethodsHpoStudyConfig {
2161 controller_id: "controller:tuner".to_string(),
2162 study_id: "study:hpo.scheduler".to_string(),
2163 methods_abi: "test-abi".to_string(),
2164 search_space: HpoSearchSpace {
2165 parameters: vec![HpoParameter::Int {
2166 name: "n_components".to_string(),
2167 low: 1,
2168 high: 1,
2169 step: 1,
2170 log: false,
2171 }],
2172 },
2173 optimizer: HpoOptimizerConfig {
2174 sampler: HpoSampler::Random,
2175 pruner: HpoPruner::None,
2176 direction: HpoDirection::Minimize,
2177 metric: HpoMetric::Rmse,
2178 seed: 13,
2179 n_startup_trials: 1,
2180 max_resource: 0,
2181 reduction_factor: 1,
2182 },
2183 },
2184 parameter_paths: BTreeMap::from([(
2185 "n_components".to_string(),
2186 "n_components".to_string(),
2187 )]),
2188 resume_checkpoint: None,
2189 resume_variants: BTreeMap::new(),
2190 resume_terminal_trials: Vec::new(),
2191 selection: RuntimeHpoSelectionTarget {
2192 producer_node: target,
2193 producer_port: "prediction".to_string(),
2194 metric: RegressionMetricKind::Rmse,
2195 direction: HpoDirection::Minimize,
2196 },
2197 provenance: RuntimeHpoProvenance {
2198 graph_fingerprint: plan.graph_fingerprint.clone(),
2199 campaign_fingerprint: plan.campaign_fingerprint.clone(),
2200 controller_fingerprint: plan.controller_fingerprint.clone(),
2201 data_identities_fingerprint: "identity:test".to_string(),
2202 fold_set_fingerprint: plan
2203 .fold_set
2204 .as_ref()
2205 .map(stable_json_fingerprint)
2206 .transpose()
2207 .unwrap(),
2208 training_influence_fingerprint: "influence:test".to_string(),
2209 relation_fingerprint: "relation:test".to_string(),
2210 },
2211 };
2212 let provider = InMemoryDataProvider::new(ControllerId::new("controller:data").unwrap());
2213 let ctx = RunContext::new(RunId::new("run:hpo.scheduler").unwrap(), Some(13));
2214
2215 let result = SequentialScheduler
2216 .execute_hpo_campaign(&plan, &controllers, &provider, &ctx, &hpo)
2217 .unwrap();
2218
2219 assert_eq!(result.operation_id, "hpo:test");
2220 assert_eq!(result.candidates.len(), 1);
2221 assert_eq!(result.checkpoint.completed_proposals.len(), 1);
2222 assert_eq!(result.checkpoint.completed_reports.len(), 1);
2223 assert_eq!(result.candidates[0].lineage.len(), 2);
2224 assert_eq!(result.incumbent.variant_id, plan.variants[0].variant_id);
2225 let mut selected_ctx = RunContext::new(RunId::new("run:hpo.scheduler").unwrap(), Some(13));
2226 selected_ctx.variant_id = Some(plan.variants[0].variant_id.clone());
2227 let selected_results = SequentialScheduler
2228 .execute_campaign_phase_with_data_provider(
2229 &plan,
2230 &controllers,
2231 &provider,
2232 &mut selected_ctx,
2233 Phase::FitCv,
2234 )
2235 .unwrap();
2236 assert_eq!(selected_results.len(), 2);
2237 assert_eq!(selected_ctx.lineage.len(), 2);
2238 assert_eq!(
2239 trace.lock().unwrap().as_slice(),
2240 [
2241 "session_factory",
2242 "ask",
2243 "model_cv",
2244 "model_cv",
2245 "intermediate",
2246 "tell",
2247 "checkpoint",
2248 "model_cv",
2249 "model_cv"
2250 ]
2251 );
2252
2253 let resumed_trace = Arc::new(Mutex::new(Vec::new()));
2257 let mut resumed_controllers = RuntimeControllerRegistry::new();
2258 resumed_controllers
2259 .register(Box::new(HpoTestTuner {
2260 id: ControllerId::new("controller:tuner").unwrap(),
2261 trace: Arc::clone(&resumed_trace),
2262 history_len: 2,
2263 proposal_count: 2,
2264 }))
2265 .unwrap();
2266 resumed_controllers
2267 .register(Box::new(HpoTestModel {
2268 id: ControllerId::new("controller:model").unwrap(),
2269 trace: Arc::clone(&resumed_trace),
2270 }))
2271 .unwrap();
2272 let mut resumed_hpo = hpo.clone();
2273 resumed_hpo.trial_budget_total = 4;
2274 let resumed_ctx = RunContext::new(RunId::new("run:hpo.resumed").unwrap(), Some(13));
2275 let resumed = SequentialScheduler
2276 .execute_hpo_campaign(
2277 &plan,
2278 &resumed_controllers,
2279 &provider,
2280 &resumed_ctx,
2281 &resumed_hpo,
2282 )
2283 .unwrap();
2284 assert_eq!(resumed.candidates.len(), 2);
2285 assert_eq!(resumed.checkpoint.trial_history_len, 4);
2286 assert_eq!(
2287 resumed_trace
2288 .lock()
2289 .unwrap()
2290 .iter()
2291 .filter(|event| event.as_str() == "ask")
2292 .count(),
2293 2
2294 );
2295
2296 let mut over_budget_controllers = RuntimeControllerRegistry::new();
2297 over_budget_controllers
2298 .register(Box::new(HpoTestTuner {
2299 id: ControllerId::new("controller:tuner").unwrap(),
2300 trace: Arc::new(Mutex::new(Vec::new())),
2301 history_len: 5,
2302 proposal_count: 0,
2303 }))
2304 .unwrap();
2305 let error = SequentialScheduler
2306 .execute_hpo_campaign(
2307 &plan,
2308 &over_budget_controllers,
2309 &provider,
2310 &resumed_ctx,
2311 &resumed_hpo,
2312 )
2313 .unwrap_err();
2314 assert!(error.to_string().contains("exceeds total trial budget"));
2315 }
2316}
2317
2318pub(crate) fn attach_coordinator_input_lineage(
2319 result: &mut NodeResult,
2320 plan: &ExecutionPlan,
2321 node_id: &NodeId,
2322 upstream_lineage: &BTreeMap<NodeId, LineageId>,
2323) -> Result<()> {
2324 let inferred = inferred_input_lineage_for_node(plan, node_id, upstream_lineage);
2325 if result.lineage.input_lineage.is_empty() {
2326 result.lineage.input_lineage = inferred;
2327 return Ok(());
2328 }
2329
2330 let declared = result
2331 .lineage
2332 .input_lineage
2333 .iter()
2334 .cloned()
2335 .collect::<BTreeSet<_>>()
2336 .into_iter()
2337 .collect::<Vec<_>>();
2338 if declared != inferred {
2339 return Err(DagMlError::RuntimeValidation(format!(
2340 "lineage for node `{}` declared input lineage {:?}, expected {:?}",
2341 result.node_id, declared, inferred
2342 )));
2343 }
2344 result.lineage.input_lineage = declared;
2345 Ok(())
2346}
2347
2348pub(crate) fn inferred_input_lineage_for_node(
2349 plan: &ExecutionPlan,
2350 node_id: &NodeId,
2351 upstream_lineage: &BTreeMap<NodeId, LineageId>,
2352) -> Vec<LineageId> {
2353 plan.graph_plan
2354 .graph
2355 .edges
2356 .iter()
2357 .filter(|edge| &edge.target.node_id == node_id && edge.contract.propagates_lineage)
2358 .filter_map(|edge| upstream_lineage.get(&edge.source.node_id).cloned())
2359 .collect::<BTreeSet<_>>()
2360 .into_iter()
2361 .collect()
2362}
2363pub(crate) fn collect_input_handles(
2364 plan: &ExecutionPlan,
2365 node_plan: &NodePlan,
2366 output_handles: &BTreeMap<NodeId, BTreeMap<String, HandleRef>>,
2367 output_data_views: &BTreeMap<NodeId, BTreeMap<String, DataProviderViewSpec>>,
2368 resources: &PhaseScopeResources<'_>,
2369 ctx: &RunContext,
2370 scope: &PhaseScope,
2371) -> Result<CollectedInputs> {
2372 let mut inputs = BTreeMap::new();
2373 let mut data_views = BTreeMap::new();
2374 let mut prediction_inputs = BTreeMap::new();
2375 let training_oof_edges = incoming_training_oof_edges(plan, node_plan, scope)?;
2376 let masked_oof_source_ports = if scope.phase == Phase::Predict {
2382 incoming_oof_edges(plan, node_plan)?
2383 } else {
2384 training_oof_edges.clone()
2385 }
2386 .into_iter()
2387 .map(|edge| (edge.source.node_id.clone(), edge.source.port_name.clone()))
2388 .collect::<BTreeSet<_>>();
2389 let bound_data_inputs = node_plan
2390 .data_bindings
2391 .iter()
2392 .map(|binding| binding.input_name.clone())
2393 .collect::<BTreeSet<_>>();
2394 let declared_source_ports = plan
2398 .graph_plan
2399 .graph
2400 .edges
2401 .iter()
2402 .filter(|edge| edge.target.node_id == node_plan.node_id)
2403 .map(|edge| (edge.source.node_id.clone(), edge.source.port_name.clone()))
2404 .collect::<BTreeSet<_>>();
2405 for upstream in &node_plan.input_nodes {
2406 if let Some(handles) = output_handles.get(upstream) {
2407 for (port, handle) in handles {
2408 if !declared_source_ports.contains(&(upstream.clone(), port.clone())) {
2409 continue;
2410 }
2411 if masked_oof_source_ports.contains(&(upstream.clone(), port.clone())) {
2412 continue;
2413 }
2414 inputs.insert(format!("{upstream}.{port}"), handle.clone());
2415 }
2416 }
2417 }
2418 for edge in plan
2419 .graph_plan
2420 .graph
2421 .edges
2422 .iter()
2423 .filter(|edge| edge.target.node_id == node_plan.node_id)
2424 .filter(|edge| edge.contract.kind == PortKind::Data && !edge.contract.requires_oof)
2425 {
2426 if bound_data_inputs.contains(&edge.target.port_name) {
2427 continue;
2428 }
2429 let Some(handles) = output_handles.get(&edge.source.node_id) else {
2430 continue;
2431 };
2432 let Some(handle) = handles.get(&edge.source.port_name) else {
2433 continue;
2434 };
2435 let key = data_view_key(&edge.target.port_name);
2436 if inputs.insert(key.clone(), handle.clone()).is_some() {
2437 return Err(DagMlError::RuntimeValidation(format!(
2438 "node `{}` received duplicate data edge input `{key}`",
2439 node_plan.node_id
2440 )));
2441 }
2442 if let Some(source_views) = output_data_views.get(&edge.source.node_id) {
2443 if let Some(view) = source_views.get(&edge.source.port_name) {
2444 if data_views.insert(key.clone(), view.clone()).is_some() {
2445 return Err(DagMlError::RuntimeValidation(format!(
2446 "node `{}` received duplicate data edge view `{key}`",
2447 node_plan.node_id
2448 )));
2449 }
2450 }
2451 let source_validation_key = validation_data_view_key(&edge.source.port_name);
2452 if let Some(view) = source_views.get(&source_validation_key) {
2453 let validation_key = format!("{key}:validation");
2454 if data_views
2455 .insert(validation_key.clone(), view.clone())
2456 .is_some()
2457 {
2458 return Err(DagMlError::RuntimeValidation(format!(
2459 "node `{}` received duplicate data edge validation view `{validation_key}`",
2460 node_plan.node_id
2461 )));
2462 }
2463 }
2464 }
2465 }
2466 for edge in training_oof_edges {
2467 let key = format!("{}.{}", edge.source.node_id, edge.source.port_name);
2468 let Some(input) = collect_oof_prediction_input(plan, edge, ctx, scope, resources)? else {
2469 return Ok(CollectedInputs {
2470 handles: BTreeMap::new(),
2471 data_views: BTreeMap::new(),
2472 prediction_inputs: BTreeMap::new(),
2473 skip_node: true,
2474 });
2475 };
2476 if inputs.insert(key.clone(), input.handle).is_some() {
2477 return Err(DagMlError::RuntimeValidation(format!(
2478 "node `{}` received duplicate OOF prediction input `{key}`",
2479 node_plan.node_id
2480 )));
2481 }
2482 if prediction_inputs.insert(key.clone(), input.spec).is_some() {
2483 return Err(DagMlError::RuntimeValidation(format!(
2484 "node `{}` received duplicate OOF prediction spec `{key}`",
2485 node_plan.node_id
2486 )));
2487 }
2488 }
2489 if matches!(scope.phase, Phase::Refit | Phase::Predict) {
2496 let off_fold_suffix = scope.phase.as_str().to_ascii_lowercase();
2497 for edge in incoming_oof_edges(plan, node_plan)? {
2498 let Some(input) = collect_off_fold_prediction_input(plan, edge, ctx, scope)? else {
2499 continue;
2500 };
2501 let key = format!(
2502 "{}.{}:{off_fold_suffix}",
2503 edge.source.node_id, edge.source.port_name
2504 );
2505 if inputs.insert(key.clone(), input.handle).is_some() {
2506 return Err(DagMlError::RuntimeValidation(format!(
2507 "node `{}` received duplicate off-fold prediction input `{key}`",
2508 node_plan.node_id
2509 )));
2510 }
2511 if prediction_inputs.insert(key.clone(), input.spec).is_some() {
2512 return Err(DagMlError::RuntimeValidation(format!(
2513 "node `{}` received duplicate off-fold prediction spec `{key}`",
2514 node_plan.node_id
2515 )));
2516 }
2517 }
2518 }
2519 if !node_plan.data_bindings.is_empty() && resources.data_provider.is_none() {
2520 return Err(DagMlError::RuntimeValidation(format!(
2521 "node `{}` requires {} data binding(s) but no runtime data provider is registered",
2522 node_plan.node_id,
2523 node_plan.data_bindings.len()
2524 )));
2525 }
2526 if let Some(data_provider) = resources.data_provider {
2527 let excluded_samples = coordinator_relations_for_node(node_plan, resources)?
2531 .map(|relations| relations.excluded_sample_ids())
2532 .unwrap_or_default();
2533 for binding in &node_plan.data_bindings {
2534 let materialized = data_provider.materialize(&DataMaterializationRequest {
2535 run_id: ctx.run_id.clone(),
2536 node_id: node_plan.node_id.clone(),
2537 input_name: binding.input_name.clone(),
2538 phase: scope.phase,
2539 variant_id: scope.variant_id.clone(),
2540 fold_id: scope.fold_id.clone(),
2541 binding: binding.clone(),
2542 })?;
2543 let branch_view_for_node = branch_view_from_node_metadata(plan, &node_plan.node_id)?;
2544 let view = data_view_for_scope(
2545 binding,
2546 plan.fold_set.as_ref(),
2547 scope,
2548 branch_view_for_node.as_ref(),
2549 &excluded_samples,
2550 )?;
2551 let key = data_view_key(&binding.input_name);
2552 let view_handle = make_data_view_handle(
2553 data_provider,
2554 ctx,
2555 node_plan,
2556 scope,
2557 binding,
2558 &materialized,
2559 &view,
2560 )?;
2561 if data_views.insert(key.clone(), view).is_some() {
2562 return Err(DagMlError::RuntimeValidation(format!(
2563 "node `{}` received duplicate data view `{key}`",
2564 node_plan.node_id
2565 )));
2566 }
2567 if inputs.insert(key.clone(), view_handle).is_some() {
2568 return Err(DagMlError::RuntimeValidation(format!(
2569 "node `{}` received duplicate data input `{key}`",
2570 node_plan.node_id
2571 )));
2572 }
2573
2574 if let Some(validation_view) = validation_data_view_for_scope(
2575 binding,
2576 plan.fold_set.as_ref(),
2577 scope,
2578 branch_view_for_node.as_ref(),
2579 &excluded_samples,
2580 )? {
2581 let validation_key = format!("{key}:validation");
2582 let validation_handle = make_data_view_handle(
2583 data_provider,
2584 ctx,
2585 node_plan,
2586 scope,
2587 binding,
2588 &materialized,
2589 &validation_view,
2590 )?;
2591 if data_views
2592 .insert(validation_key.clone(), validation_view)
2593 .is_some()
2594 {
2595 return Err(DagMlError::RuntimeValidation(format!(
2596 "node `{}` received duplicate validation data view `{validation_key}`",
2597 node_plan.node_id
2598 )));
2599 }
2600 if inputs
2601 .insert(validation_key.clone(), validation_handle)
2602 .is_some()
2603 {
2604 return Err(DagMlError::RuntimeValidation(format!(
2605 "node `{}` received duplicate validation data input `{validation_key}`",
2606 node_plan.node_id
2607 )));
2608 }
2609 }
2610 }
2611 }
2612 Ok(CollectedInputs {
2613 handles: inputs,
2614 data_views,
2615 prediction_inputs,
2616 skip_node: false,
2617 })
2618}
2619pub(crate) fn preload_replay_prediction_cache_store(
2620 bundle: &ExecutionBundle,
2621 prediction_cache_store: Option<&dyn RuntimePredictionCacheStore>,
2622 ctx: &mut RunContext,
2623) -> Result<()> {
2624 if bundle.prediction_requirements.is_empty() {
2625 return Ok(());
2626 }
2627 let store = prediction_cache_store.ok_or_else(|| {
2628 DagMlError::RuntimeValidation(format!(
2629 "bundle `{}` cannot preload OOF prediction caches without a prediction cache store",
2630 bundle.bundle_id
2631 ))
2632 })?;
2633 if !ctx.prediction_store.blocks().is_empty() {
2634 return Err(DagMlError::RuntimeValidation(format!(
2635 "bundle `{}` cannot preload OOF prediction caches into a non-empty prediction store",
2636 bundle.bundle_id
2637 )));
2638 }
2639 let contracts = replay_prediction_cache_contracts(bundle)?;
2640 for contract in contracts.values() {
2641 if contract.requirement.prediction_level == PredictionLevel::Sample {
2642 let blocks = store.load_blocks(&contract.cache.requirement_key)?;
2643 if blocks.iter().any(|block| {
2644 block.producer_node != contract.requirement.producer_node
2645 || block.partition != contract.requirement.partition
2646 }) {
2647 return Err(DagMlError::RuntimeValidation(format!(
2648 "prediction cache store returned blocks outside requirement `{}`",
2649 contract.cache.requirement_key
2650 )));
2651 }
2652 let mut payload = build_prediction_cache_payload(&contract.requirement, &blocks)?;
2653 payload.cache_namespace_fingerprints =
2654 contract.cache.cache_namespace_fingerprints.clone();
2655 validate_prediction_cache_payload_matches_record(&payload, &contract.cache)?;
2656 for block in &payload.blocks {
2657 ctx.prediction_store.append(block.clone())?;
2658 }
2659 } else {
2660 let blocks = store.load_aggregated_blocks(&contract.cache.requirement_key)?;
2661 if blocks.iter().any(|block| {
2662 block.producer_node != contract.requirement.producer_node
2663 || block.partition != contract.requirement.partition
2664 || block.level != contract.requirement.prediction_level
2665 }) {
2666 return Err(DagMlError::RuntimeValidation(format!(
2667 "prediction cache store returned aggregated blocks outside requirement `{}`",
2668 contract.cache.requirement_key
2669 )));
2670 }
2671 let mut payload =
2672 build_aggregated_prediction_cache_payload(&contract.requirement, &blocks)?;
2673 payload.cache_namespace_fingerprints =
2674 contract.cache.cache_namespace_fingerprints.clone();
2675 validate_prediction_cache_payload_matches_record(&payload, &contract.cache)?;
2676 }
2677 }
2678 Ok(())
2679}
2680
2681pub(crate) fn replay_prediction_cache_contracts(
2682 bundle: &ExecutionBundle,
2683) -> Result<BTreeMap<String, ReplayPredictionCacheContract>> {
2684 bundle.validate()?;
2685 let requirements = bundle
2686 .prediction_requirements
2687 .iter()
2688 .map(|requirement| (requirement.key(), requirement))
2689 .collect::<BTreeMap<_, _>>();
2690 let mut contracts = BTreeMap::new();
2691 for cache in &bundle.prediction_caches {
2692 let requirement = requirements.get(&cache.requirement_key).ok_or_else(|| {
2693 DagMlError::RuntimeValidation(format!(
2694 "prediction cache `{}` references unknown prediction requirement `{}`",
2695 cache.cache_id, cache.requirement_key
2696 ))
2697 })?;
2698 contracts.insert(
2699 cache.requirement_key.clone(),
2700 ReplayPredictionCacheContract {
2701 requirement: (*requirement).clone(),
2702 cache: cache.clone(),
2703 },
2704 );
2705 }
2706 Ok(contracts)
2707}
2708
2709pub(crate) fn materialize_replay_artifact_handles(
2710 plan: &ExecutionPlan,
2711 bundle: &ExecutionBundle,
2712 replay_request: &ReplayPhaseRequest,
2713 artifact_store: &dyn RuntimeArtifactStore,
2714 ctx: &RunContext,
2715) -> Result<MaterializedReplayArtifacts> {
2716 let mut handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
2717 let mut inputs = BTreeMap::<NodeId, BTreeMap<String, ArtifactInputSpec>>::new();
2718 for artifact in &bundle.refit_artifacts {
2719 artifact.validate()?;
2720 let node_plan = plan.node_plans.get(&artifact.node_id).ok_or_else(|| {
2721 DagMlError::RuntimeValidation(format!(
2722 "bundle `{}` artifact references unknown node `{}`",
2723 bundle.bundle_id, artifact.node_id
2724 ))
2725 })?;
2726 if !node_plan.supported_phases.contains(&replay_request.phase) {
2727 return Err(DagMlError::RuntimeValidation(format!(
2728 "bundle `{}` artifact node `{}` does not support replay phase {:?}",
2729 bundle.bundle_id, artifact.node_id, replay_request.phase
2730 )));
2731 }
2732 let handle = artifact_store.materialize(&ArtifactMaterializationRequest {
2733 run_id: ctx.run_id.clone(),
2734 bundle_id: bundle.bundle_id.clone(),
2735 node_id: artifact.node_id.clone(),
2736 phase: replay_request.phase,
2737 variant_id: bundle.selected_variant_id.clone(),
2738 controller_id: artifact.controller_id.clone(),
2739 artifact: artifact.artifact.clone(),
2740 params_fingerprint: artifact.params_fingerprint.clone(),
2741 training_loss_fingerprint: artifact.training_loss_fingerprint.clone(),
2742 })?;
2743 if !matches!(handle.kind, HandleKind::Model | HandleKind::Artifact) {
2744 return Err(DagMlError::RuntimeValidation(format!(
2745 "artifact `{}` materialized as unsupported handle kind {:?}",
2746 artifact.artifact.id, handle.kind
2747 )));
2748 }
2749 if handle.owner_controller != artifact.controller_id {
2750 return Err(DagMlError::RuntimeValidation(format!(
2751 "artifact `{}` handle owner `{}` does not match controller `{}`",
2752 artifact.artifact.id, handle.owner_controller, artifact.controller_id
2753 )));
2754 }
2755 let key = refit_artifact_input_key(&artifact.artifact.id);
2756 if handles
2757 .entry(artifact.node_id.clone())
2758 .or_default()
2759 .insert(key.clone(), handle)
2760 .is_some()
2761 {
2762 return Err(DagMlError::RuntimeValidation(format!(
2763 "duplicate replay artifact input `{key}` for node `{}`",
2764 artifact.node_id
2765 )));
2766 }
2767 if inputs
2768 .entry(artifact.node_id.clone())
2769 .or_default()
2770 .insert(key.clone(), ArtifactInputSpec::from_refit_record(artifact)?)
2771 .is_some()
2772 {
2773 return Err(DagMlError::RuntimeValidation(format!(
2774 "duplicate replay artifact metadata `{key}` for node `{}`",
2775 artifact.node_id
2776 )));
2777 }
2778 }
2779 Ok(MaterializedReplayArtifacts { handles, inputs })
2780}
2781
2782pub(crate) fn derive_task_seed(
2783 root_seed: Option<u64>,
2784 variant_id: Option<&VariantId>,
2785 fold_id: Option<&FoldId>,
2786 node_plan: &NodePlan,
2787 phase: Phase,
2788) -> Option<u64> {
2789 root_seed.map(|root| {
2790 let mut context = SeedContext::root(root);
2791 if let Some(variant_id) = variant_id {
2792 context = context.child(format!("variant:{variant_id}"));
2793 }
2794 if let Some(fold_id) = fold_id {
2795 context = context.child(format!("fold:{fold_id}"));
2796 }
2797 context
2798 .child(format!("node:{}", node_plan.node_id))
2799 .child(format!("phase:{phase:?}"))
2800 .derive_u64("task")
2801 })
2802}