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