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_phase(
170 &self,
171 plan: &ExecutionPlan,
172 controllers: &RuntimeControllerRegistry,
173 ctx: &mut RunContext,
174 phase: Phase,
175 ) -> Result<Vec<NodeResult>> {
176 plan.validate()?;
177 let variant_id = ctx.variant_id.clone();
178 let seed_root = ctx.root_seed;
179 self.execute_phase_scope(
180 plan,
181 controllers,
182 ctx,
183 PhaseScope {
184 phase,
185 variant_id,
186 variant: None,
187 fold_id: None,
188 seed_root,
189 },
190 PhaseScopeResources::default(),
191 )
192 }
193
194 pub fn execute_phase_with_data_provider(
195 &self,
196 plan: &ExecutionPlan,
197 controllers: &RuntimeControllerRegistry,
198 data_provider: &dyn RuntimeDataProvider,
199 ctx: &mut RunContext,
200 phase: Phase,
201 ) -> Result<Vec<NodeResult>> {
202 plan.validate()?;
203 let variant_id = ctx.variant_id.clone();
204 let seed_root = ctx.root_seed;
205 self.execute_phase_scope(
206 plan,
207 controllers,
208 ctx,
209 PhaseScope {
210 phase,
211 variant_id,
212 variant: None,
213 fold_id: None,
214 seed_root,
215 },
216 PhaseScopeResources {
217 data_provider: Some(data_provider),
218 ..Default::default()
219 },
220 )
221 }
222
223 pub fn execute_campaign_phase(
224 &self,
225 plan: &ExecutionPlan,
226 controllers: &RuntimeControllerRegistry,
227 ctx: &mut RunContext,
228 phase: Phase,
229 ) -> Result<Vec<NodeResult>> {
230 plan.validate()?;
231 let mut results = Vec::new();
232 let fold_ids = if phase == Phase::FitCv {
233 plan.fold_set
234 .as_ref()
235 .map(|fold_set| {
236 fold_set
237 .folds
238 .iter()
239 .map(|fold| Some(fold.fold_id.clone()))
240 .collect::<Vec<_>>()
241 })
242 .unwrap_or_else(|| vec![None])
243 } else {
244 vec![None]
245 };
246 for variant in &plan.variants {
247 if ctx
248 .variant_id
249 .as_ref()
250 .is_some_and(|requested| requested != &variant.variant_id)
251 {
252 continue;
253 }
254 for fold_id in &fold_ids {
255 let seed_root = variant.seed.or(ctx.root_seed);
256 results.extend(self.execute_phase_scope(
257 plan,
258 controllers,
259 ctx,
260 PhaseScope {
261 phase,
262 variant_id: Some(variant.variant_id.clone()),
263 variant: Some(VariantExecutionSpec::from_plan(variant)),
264 fold_id: fold_id.clone(),
265 seed_root,
266 },
267 PhaseScopeResources::default(),
268 )?);
269 }
270 }
271 Ok(results)
272 }
273
274 pub fn execute_campaign_phase_with_data_provider(
275 &self,
276 plan: &ExecutionPlan,
277 controllers: &RuntimeControllerRegistry,
278 data_provider: &dyn RuntimeDataProvider,
279 ctx: &mut RunContext,
280 phase: Phase,
281 ) -> Result<Vec<NodeResult>> {
282 plan.validate()?;
283 let mut results = Vec::new();
284 let fold_ids = if phase == Phase::FitCv {
285 plan.fold_set
286 .as_ref()
287 .map(|fold_set| {
288 fold_set
289 .folds
290 .iter()
291 .map(|fold| Some(fold.fold_id.clone()))
292 .collect::<Vec<_>>()
293 })
294 .unwrap_or_else(|| vec![None])
295 } else {
296 vec![None]
297 };
298 for variant in &plan.variants {
299 if ctx
300 .variant_id
301 .as_ref()
302 .is_some_and(|requested| requested != &variant.variant_id)
303 {
304 continue;
305 }
306 for fold_id in &fold_ids {
307 let seed_root = variant.seed.or(ctx.root_seed);
308 results.extend(self.execute_phase_scope(
309 plan,
310 controllers,
311 ctx,
312 PhaseScope {
313 phase,
314 variant_id: Some(variant.variant_id.clone()),
315 variant: Some(VariantExecutionSpec::from_plan(variant)),
316 fold_id: fold_id.clone(),
317 seed_root,
318 },
319 PhaseScopeResources {
320 data_provider: Some(data_provider),
321 ..Default::default()
322 },
323 )?);
324 }
325 }
326 Ok(results)
327 }
328
329 pub fn execute_campaign_phase_with_data_provider_and_artifact_store(
330 &self,
331 plan: &ExecutionPlan,
332 controllers: &RuntimeControllerRegistry,
333 data_provider: &dyn RuntimeDataProvider,
334 artifact_store: &mut InMemoryArtifactStore,
335 ctx: &mut RunContext,
336 phase: Phase,
337 ) -> Result<Vec<NodeResult>> {
338 plan.validate()?;
339 let mut results = Vec::new();
340 let fold_ids = if phase == Phase::FitCv {
341 plan.fold_set
342 .as_ref()
343 .map(|fold_set| {
344 fold_set
345 .folds
346 .iter()
347 .map(|fold| Some(fold.fold_id.clone()))
348 .collect::<Vec<_>>()
349 })
350 .unwrap_or_else(|| vec![None])
351 } else {
352 vec![None]
353 };
354 for variant in &plan.variants {
355 if ctx
356 .variant_id
357 .as_ref()
358 .is_some_and(|requested| requested != &variant.variant_id)
359 {
360 continue;
361 }
362 for fold_id in &fold_ids {
363 let seed_root = variant.seed.or(ctx.root_seed);
364 results.extend(self.execute_phase_scope(
365 plan,
366 controllers,
367 ctx,
368 PhaseScope {
369 phase,
370 variant_id: Some(variant.variant_id.clone()),
371 variant: Some(VariantExecutionSpec::from_plan(variant)),
372 fold_id: fold_id.clone(),
373 seed_root,
374 },
375 PhaseScopeResources {
376 data_provider: Some(data_provider),
377 artifact_store: Some(&mut *artifact_store),
378 ..Default::default()
379 },
380 )?);
381 }
382 }
383 Ok(results)
384 }
385
386 pub fn execute_bundle_replay(
387 &self,
388 replay: BundleReplayExecution<'_>,
389 ctx: &mut RunContext,
390 ) -> Result<Vec<NodeResult>> {
391 replay.bundle.validate_against_plan(replay.plan)?;
392 replay
393 .replay_request
394 .validate_for_bundle_with_prediction_cache_store(
395 replay.bundle,
396 replay.prediction_cache_store.is_some(),
397 )?;
398 replay
399 .bundle
400 .validate_replay_envelopes(replay.data_envelopes)?;
401 let prediction_cache_contracts = if replay.replay_request.phase == Phase::Refit {
402 Some(replay_prediction_cache_contracts(replay.bundle)?)
403 } else {
404 None
405 };
406 if replay.replay_request.phase == Phase::Refit {
407 preload_replay_prediction_cache_store(
408 replay.bundle,
409 replay.prediction_cache_store,
410 ctx,
411 )?;
412 }
413 let replay_artifacts = materialize_replay_artifact_handles(
414 replay.plan,
415 replay.bundle,
416 replay.replay_request,
417 replay.artifact_store,
418 ctx,
419 )?;
420 let selected_variant = replay
421 .bundle
422 .selected_variant_id
423 .as_ref()
424 .map(|selected| {
425 replay
426 .plan
427 .variants
428 .iter()
429 .find(|variant| &variant.variant_id == selected)
430 .map(VariantExecutionSpec::from_plan)
431 .ok_or_else(|| {
432 DagMlError::RuntimeValidation(format!(
433 "bundle `{}` selected unknown variant `{selected}`",
434 replay.bundle.bundle_id
435 ))
436 })
437 })
438 .transpose()?;
439 let seed_root = selected_variant
440 .as_ref()
441 .and_then(|variant| variant.seed)
442 .or(ctx.root_seed);
443
444 self.execute_phase_scope(
445 replay.plan,
446 replay.controllers,
447 ctx,
448 PhaseScope {
449 phase: replay.replay_request.phase,
450 variant_id: replay.bundle.selected_variant_id.clone(),
451 variant: selected_variant,
452 fold_id: None,
453 seed_root,
454 },
455 PhaseScopeResources {
456 data_provider: Some(replay.data_provider),
457 replay_artifact_handles: Some(&replay_artifacts.handles),
458 replay_artifact_inputs: Some(&replay_artifacts.inputs),
459 replay_bundle_id: Some(&replay.bundle.bundle_id),
460 data_envelopes: Some(replay.data_envelopes),
461 prediction_cache_store: replay.prediction_cache_store,
462 prediction_cache_contracts: prediction_cache_contracts.as_ref(),
463 ..Default::default()
464 },
465 )
466 }
467
468 fn execute_phase_scope(
469 &self,
470 plan: &ExecutionPlan,
471 controllers: &RuntimeControllerRegistry,
472 ctx: &mut RunContext,
473 scope: PhaseScope,
474 mut resources: PhaseScopeResources<'_>,
475 ) -> Result<Vec<NodeResult>> {
476 let _phase_span = crate::observability::phase_span(
477 ctx.run_id.as_str(),
478 plan.id.as_str(),
479 scope.phase.as_str(),
480 scope.variant_id.as_ref().map(VariantId::as_str),
481 scope.fold_id.as_ref().map(FoldId::as_str),
482 )
483 .entered();
484 let mut results = Vec::new();
485 let mut output_handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
486 let mut output_data_views =
487 BTreeMap::<NodeId, BTreeMap<String, DataProviderViewSpec>>::new();
488 let mut input_lineage = BTreeMap::<NodeId, LineageId>::new();
489
490 for level in plan.node_parallel_levels_for_phase(scope.phase)? {
491 for node_id in &level {
492 let node_plan = plan
493 .node_plans
494 .get(node_id)
495 .expect("execution plan was validated");
496 if let Some(reduction) = merge_reduction_mode(plan, node_plan) {
503 if let Some(mut result) =
504 reassemble_branch_merge(plan, node_plan, ctx, &scope, reduction)?
505 {
506 let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
507 let task = NodeTask {
508 inner_fold_set: None,
509 run_id: ctx.run_id.clone(),
510 node_plan: task_node_plan.clone(),
511 phase: scope.phase,
512 variant_id: scope.variant_id.clone(),
513 variant: scope.variant.clone(),
514 fold_id: scope.fold_id.clone(),
515 branch_path: Vec::new(),
516 input_handles: BTreeMap::new(),
517 data_views: BTreeMap::new(),
518 prediction_inputs: BTreeMap::new(),
519 artifact_inputs: BTreeMap::new(),
520 required_loss_attestations: NodeTask::required_loss_attestations_for(
521 &task_node_plan,
522 scope.phase,
523 )?,
524 fit_influence: FitInfluenceTask::default(),
525 seed: None,
526 };
527 normalize_result_prediction_ports(plan, &task, &mut result)?;
528 result.validate_for_task(&task)?;
529 for prediction in &result.predictions {
530 ctx.prediction_store.append(prediction.clone())?;
531 }
532 apply_result_scoring(
533 &result,
534 &mut ctx.score_collector,
535 &mut ctx.regression_target_records,
536 )?;
537 ctx.lineage.record(result.lineage.clone())?;
538 output_handles.insert(node_id.clone(), result.outputs.clone());
539 input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
540 results.push(result);
541 }
542 continue;
543 }
544 let controller = controllers.get(&node_plan.controller_id).ok_or_else(|| {
545 DagMlError::RuntimeValidation(format!(
546 "runtime controller `{}` is not registered",
547 node_plan.controller_id
548 ))
549 })?;
550 let collected_inputs = collect_input_handles(
551 plan,
552 node_plan,
553 &output_handles,
554 &output_data_views,
555 &resources,
556 ctx,
557 &scope,
558 )?;
559 if collected_inputs.skip_node {
560 continue;
561 }
562 let mut input_handles = collected_inputs.handles;
563 let mut artifact_inputs = BTreeMap::new();
564 if let Some(node_artifact_handles) = resources
565 .replay_artifact_handles
566 .and_then(|handles| handles.get(node_id))
567 {
568 for (key, handle) in node_artifact_handles {
569 if input_handles.insert(key.clone(), handle.clone()).is_some() {
570 return Err(DagMlError::RuntimeValidation(format!(
571 "node `{node_id}` received duplicate replay artifact input `{key}`"
572 )));
573 }
574 }
575 }
576 if let Some(node_artifact_inputs) = resources
577 .replay_artifact_inputs
578 .and_then(|inputs| inputs.get(node_id))
579 {
580 for (key, spec) in node_artifact_inputs {
581 if artifact_inputs.insert(key.clone(), spec.clone()).is_some() {
582 return Err(DagMlError::RuntimeValidation(format!(
583 "node `{node_id}` received duplicate replay artifact metadata `{key}`"
584 )));
585 }
586 }
587 }
588 let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
589 let inner_fold_set = inner_fold_set_for_scope(
590 &plan.campaign,
591 plan.fold_set.as_ref(),
592 node_plan,
593 &scope,
594 )?;
595 let fit_influence = fit_influence_task_for_node(
596 plan,
597 &task_node_plan,
598 &collected_inputs.data_views,
599 )?;
600 let task = NodeTask {
601 inner_fold_set,
602 run_id: ctx.run_id.clone(),
603 node_plan: task_node_plan.clone(),
604 phase: scope.phase,
605 variant_id: scope.variant_id.clone(),
606 variant: scope.variant.clone(),
607 fold_id: scope.fold_id.clone(),
608 branch_path: Vec::new(),
609 input_handles,
610 data_views: collected_inputs.data_views,
611 prediction_inputs: collected_inputs.prediction_inputs,
612 artifact_inputs,
613 required_loss_attestations: NodeTask::required_loss_attestations_for(
614 &task_node_plan,
615 scope.phase,
616 )?,
617 fit_influence,
618 seed: derive_task_seed(
619 scope.seed_root,
620 scope.variant_id.as_ref(),
621 scope.fold_id.as_ref(),
622 &task_node_plan,
623 scope.phase,
624 ),
625 };
626 let _node_span = crate::observability::node_span(
627 task.run_id.as_str(),
628 plan.id.as_str(),
629 task.phase.as_str(),
630 task.node_plan.node_id.as_str(),
631 task.node_plan.controller_id.as_str(),
632 )
633 .entered();
634 let mut result = controller.invoke(&task)?;
635 record_fit_influence_diagnostic(&task, &mut result);
636 normalize_result_prediction_ports(plan, &task, &mut result)?;
637 result.validate_for_task(&task)?;
638 apply_result_prediction_aggregation(
639 plan,
640 controllers,
641 &task,
642 &mut result,
643 &resources,
644 )?;
645 attach_coordinator_input_lineage(
646 &mut result,
647 plan,
648 &task.node_plan.node_id,
649 &input_lineage,
650 )?;
651 if let Some(store) = resources.artifact_store.as_deref_mut() {
652 if scope.phase == Phase::Refit {
653 store.capture_refit_artifacts(&task, &result)?;
654 }
655 }
656 for prediction in &result.predictions {
657 ctx.prediction_store.append(prediction.clone())?;
658 }
659 for prediction in &result.aggregated_predictions {
660 ctx.aggregated_prediction_store.append(prediction.clone())?;
661 }
662 apply_result_scoring(
663 &result,
664 &mut ctx.score_collector,
665 &mut ctx.regression_target_records,
666 )?;
667 ctx.lineage.record(result.lineage.clone())?;
668 let data_views = derive_output_data_views(plan, &task, &result)?;
669 output_handles.insert(node_id.clone(), result.outputs.clone());
670 output_data_views.insert(node_id.clone(), data_views);
671 input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
672 results.push(result);
673 }
674 }
675
676 Ok(results)
677 }
678}
679
680impl ParallelScheduler {
681 pub fn execute_phase(
682 &self,
683 plan: &ExecutionPlan,
684 controllers: &RuntimeControllerRegistry,
685 ctx: &mut RunContext,
686 phase: Phase,
687 ) -> Result<Vec<NodeResult>> {
688 plan.validate()?;
689 let variant_id = ctx.variant_id.clone();
690 let seed_root = ctx.root_seed;
691 self.execute_phase_scope(
692 plan,
693 controllers,
694 ctx,
695 PhaseScope {
696 phase,
697 variant_id,
698 variant: None,
699 fold_id: None,
700 seed_root,
701 },
702 PhaseScopeResources::default(),
703 )
704 }
705
706 pub fn execute_phase_with_data_provider(
707 &self,
708 plan: &ExecutionPlan,
709 controllers: &RuntimeControllerRegistry,
710 data_provider: &dyn RuntimeDataProvider,
711 ctx: &mut RunContext,
712 phase: Phase,
713 ) -> Result<Vec<NodeResult>> {
714 plan.validate()?;
715 let variant_id = ctx.variant_id.clone();
716 let seed_root = ctx.root_seed;
717 self.execute_phase_scope(
718 plan,
719 controllers,
720 ctx,
721 PhaseScope {
722 phase,
723 variant_id,
724 variant: None,
725 fold_id: None,
726 seed_root,
727 },
728 PhaseScopeResources {
729 data_provider: Some(data_provider),
730 ..Default::default()
731 },
732 )
733 }
734
735 pub fn execute_campaign_phase(
736 &self,
737 plan: &ExecutionPlan,
738 controllers: &RuntimeControllerRegistry,
739 ctx: &mut RunContext,
740 phase: Phase,
741 ) -> Result<Vec<NodeResult>> {
742 plan.validate()?;
743 let mut results = Vec::new();
744 let fold_ids = if phase == Phase::FitCv {
745 plan.fold_set
746 .as_ref()
747 .map(|fold_set| {
748 fold_set
749 .folds
750 .iter()
751 .map(|fold| Some(fold.fold_id.clone()))
752 .collect::<Vec<_>>()
753 })
754 .unwrap_or_else(|| vec![None])
755 } else {
756 vec![None]
757 };
758 for variant in &plan.variants {
759 if ctx
760 .variant_id
761 .as_ref()
762 .is_some_and(|requested| requested != &variant.variant_id)
763 {
764 continue;
765 }
766 for fold_id in &fold_ids {
767 let seed_root = variant.seed.or(ctx.root_seed);
768 results.extend(self.execute_phase_scope(
769 plan,
770 controllers,
771 ctx,
772 PhaseScope {
773 phase,
774 variant_id: Some(variant.variant_id.clone()),
775 variant: Some(VariantExecutionSpec::from_plan(variant)),
776 fold_id: fold_id.clone(),
777 seed_root,
778 },
779 PhaseScopeResources::default(),
780 )?);
781 }
782 }
783 Ok(results)
784 }
785
786 pub fn execute_campaign_phase_with_data_provider(
787 &self,
788 plan: &ExecutionPlan,
789 controllers: &RuntimeControllerRegistry,
790 data_provider: &dyn RuntimeDataProvider,
791 ctx: &mut RunContext,
792 phase: Phase,
793 ) -> Result<Vec<NodeResult>> {
794 plan.validate()?;
795 let mut results = Vec::new();
796 let fold_ids = if phase == Phase::FitCv {
797 plan.fold_set
798 .as_ref()
799 .map(|fold_set| {
800 fold_set
801 .folds
802 .iter()
803 .map(|fold| Some(fold.fold_id.clone()))
804 .collect::<Vec<_>>()
805 })
806 .unwrap_or_else(|| vec![None])
807 } else {
808 vec![None]
809 };
810 for variant in &plan.variants {
811 if ctx
812 .variant_id
813 .as_ref()
814 .is_some_and(|requested| requested != &variant.variant_id)
815 {
816 continue;
817 }
818 for fold_id in &fold_ids {
819 let seed_root = variant.seed.or(ctx.root_seed);
820 results.extend(self.execute_phase_scope(
821 plan,
822 controllers,
823 ctx,
824 PhaseScope {
825 phase,
826 variant_id: Some(variant.variant_id.clone()),
827 variant: Some(VariantExecutionSpec::from_plan(variant)),
828 fold_id: fold_id.clone(),
829 seed_root,
830 },
831 PhaseScopeResources {
832 data_provider: Some(data_provider),
833 ..Default::default()
834 },
835 )?);
836 }
837 }
838 Ok(results)
839 }
840
841 pub fn execute_campaign_phase_with_data_provider_and_artifact_store(
842 &self,
843 plan: &ExecutionPlan,
844 controllers: &RuntimeControllerRegistry,
845 data_provider: &dyn RuntimeDataProvider,
846 artifact_store: &mut InMemoryArtifactStore,
847 ctx: &mut RunContext,
848 phase: Phase,
849 ) -> Result<Vec<NodeResult>> {
850 plan.validate()?;
851 let mut results = Vec::new();
852 let fold_ids = if phase == Phase::FitCv {
853 plan.fold_set
854 .as_ref()
855 .map(|fold_set| {
856 fold_set
857 .folds
858 .iter()
859 .map(|fold| Some(fold.fold_id.clone()))
860 .collect::<Vec<_>>()
861 })
862 .unwrap_or_else(|| vec![None])
863 } else {
864 vec![None]
865 };
866 for variant in &plan.variants {
867 if ctx
868 .variant_id
869 .as_ref()
870 .is_some_and(|requested| requested != &variant.variant_id)
871 {
872 continue;
873 }
874 for fold_id in &fold_ids {
875 let seed_root = variant.seed.or(ctx.root_seed);
876 results.extend(self.execute_phase_scope(
877 plan,
878 controllers,
879 ctx,
880 PhaseScope {
881 phase,
882 variant_id: Some(variant.variant_id.clone()),
883 variant: Some(VariantExecutionSpec::from_plan(variant)),
884 fold_id: fold_id.clone(),
885 seed_root,
886 },
887 PhaseScopeResources {
888 data_provider: Some(data_provider),
889 artifact_store: Some(&mut *artifact_store),
890 ..Default::default()
891 },
892 )?);
893 }
894 }
895 Ok(results)
896 }
897
898 pub fn execute_bundle_replay(
899 &self,
900 replay: BundleReplayExecution<'_>,
901 ctx: &mut RunContext,
902 ) -> Result<Vec<NodeResult>> {
903 replay.bundle.validate_against_plan(replay.plan)?;
904 replay
905 .replay_request
906 .validate_for_bundle_with_prediction_cache_store(
907 replay.bundle,
908 replay.prediction_cache_store.is_some(),
909 )?;
910 replay
911 .bundle
912 .validate_replay_envelopes(replay.data_envelopes)?;
913 let prediction_cache_contracts = if replay.replay_request.phase == Phase::Refit {
914 Some(replay_prediction_cache_contracts(replay.bundle)?)
915 } else {
916 None
917 };
918 if replay.replay_request.phase == Phase::Refit {
919 preload_replay_prediction_cache_store(
920 replay.bundle,
921 replay.prediction_cache_store,
922 ctx,
923 )?;
924 }
925 let replay_artifacts = materialize_replay_artifact_handles(
926 replay.plan,
927 replay.bundle,
928 replay.replay_request,
929 replay.artifact_store,
930 ctx,
931 )?;
932 let selected_variant = replay
933 .bundle
934 .selected_variant_id
935 .as_ref()
936 .map(|selected| {
937 replay
938 .plan
939 .variants
940 .iter()
941 .find(|variant| &variant.variant_id == selected)
942 .map(VariantExecutionSpec::from_plan)
943 .ok_or_else(|| {
944 DagMlError::RuntimeValidation(format!(
945 "bundle `{}` selected unknown variant `{selected}`",
946 replay.bundle.bundle_id
947 ))
948 })
949 })
950 .transpose()?;
951 let seed_root = selected_variant
952 .as_ref()
953 .and_then(|variant| variant.seed)
954 .or(ctx.root_seed);
955
956 self.execute_phase_scope(
957 replay.plan,
958 replay.controllers,
959 ctx,
960 PhaseScope {
961 phase: replay.replay_request.phase,
962 variant_id: replay.bundle.selected_variant_id.clone(),
963 variant: selected_variant,
964 fold_id: None,
965 seed_root,
966 },
967 PhaseScopeResources {
968 data_provider: Some(replay.data_provider),
969 replay_artifact_handles: Some(&replay_artifacts.handles),
970 replay_artifact_inputs: Some(&replay_artifacts.inputs),
971 replay_bundle_id: Some(&replay.bundle.bundle_id),
972 data_envelopes: Some(replay.data_envelopes),
973 prediction_cache_store: replay.prediction_cache_store,
974 prediction_cache_contracts: prediction_cache_contracts.as_ref(),
975 ..Default::default()
976 },
977 )
978 }
979
980 fn execute_phase_scope(
981 &self,
982 plan: &ExecutionPlan,
983 controllers: &RuntimeControllerRegistry,
984 ctx: &mut RunContext,
985 scope: PhaseScope,
986 mut resources: PhaseScopeResources<'_>,
987 ) -> Result<Vec<NodeResult>> {
988 let phase_span = crate::observability::phase_span(
992 ctx.run_id.as_str(),
993 plan.id.as_str(),
994 scope.phase.as_str(),
995 scope.variant_id.as_ref().map(VariantId::as_str),
996 scope.fold_id.as_ref().map(FoldId::as_str),
997 );
998 let _phase_entered = phase_span.clone().entered();
999 let plan_id = plan.id.as_str();
1001 plan.validate_parallel_controller_capabilities(self.max_workers, scope.phase)?;
1002 let mut results = Vec::new();
1003 let mut output_handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
1004 let mut output_data_views =
1005 BTreeMap::<NodeId, BTreeMap<String, DataProviderViewSpec>>::new();
1006 let mut input_lineage = BTreeMap::<NodeId, LineageId>::new();
1007
1008 for level in plan.node_parallel_levels_for_phase(scope.phase)? {
1009 let mut prepared = Vec::<PreparedNodeTask>::new();
1010 let mut merge_nodes = Vec::<(NodeId, MergeReduction)>::new();
1017 for node_id in &level {
1018 let node_plan = plan
1019 .node_plans
1020 .get(node_id)
1021 .expect("execution plan was validated");
1022 if let Some(reduction) = merge_reduction_mode(plan, node_plan) {
1023 merge_nodes.push((node_id.clone(), reduction));
1024 continue;
1025 }
1026 let collected_inputs = collect_input_handles(
1027 plan,
1028 node_plan,
1029 &output_handles,
1030 &output_data_views,
1031 &resources,
1032 ctx,
1033 &scope,
1034 )?;
1035 if collected_inputs.skip_node {
1036 continue;
1037 }
1038 let mut input_handles = collected_inputs.handles;
1039 let mut artifact_inputs = BTreeMap::new();
1040 if let Some(node_artifact_handles) = resources
1041 .replay_artifact_handles
1042 .and_then(|handles| handles.get(node_id))
1043 {
1044 for (key, handle) in node_artifact_handles {
1045 if input_handles.insert(key.clone(), handle.clone()).is_some() {
1046 return Err(DagMlError::RuntimeValidation(format!(
1047 "node `{node_id}` received duplicate replay artifact input `{key}`"
1048 )));
1049 }
1050 }
1051 }
1052 if let Some(node_artifact_inputs) = resources
1053 .replay_artifact_inputs
1054 .and_then(|inputs| inputs.get(node_id))
1055 {
1056 for (key, spec) in node_artifact_inputs {
1057 if artifact_inputs.insert(key.clone(), spec.clone()).is_some() {
1058 return Err(DagMlError::RuntimeValidation(format!(
1059 "node `{node_id}` received duplicate replay artifact metadata `{key}`"
1060 )));
1061 }
1062 }
1063 }
1064 let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
1065 let inner_fold_set = inner_fold_set_for_scope(
1066 &plan.campaign,
1067 plan.fold_set.as_ref(),
1068 node_plan,
1069 &scope,
1070 )?;
1071 let fit_influence = fit_influence_task_for_node(
1072 plan,
1073 &task_node_plan,
1074 &collected_inputs.data_views,
1075 )?;
1076 prepared.push(PreparedNodeTask {
1077 node_id: node_id.clone(),
1078 task: NodeTask {
1079 inner_fold_set,
1080 run_id: ctx.run_id.clone(),
1081 node_plan: task_node_plan.clone(),
1082 phase: scope.phase,
1083 variant_id: scope.variant_id.clone(),
1084 variant: scope.variant.clone(),
1085 fold_id: scope.fold_id.clone(),
1086 branch_path: Vec::new(),
1087 input_handles,
1088 data_views: collected_inputs.data_views,
1089 prediction_inputs: collected_inputs.prediction_inputs,
1090 artifact_inputs,
1091 required_loss_attestations: NodeTask::required_loss_attestations_for(
1092 &task_node_plan,
1093 scope.phase,
1094 )?,
1095 fit_influence,
1096 seed: derive_task_seed(
1097 scope.seed_root,
1098 scope.variant_id.as_ref(),
1099 scope.fold_id.as_ref(),
1100 &task_node_plan,
1101 scope.phase,
1102 ),
1103 },
1104 });
1105 }
1106
1107 for chunk in prepared.chunks(self.max_workers) {
1108 let chunk_results =
1109 std::thread::scope(|thread_scope| -> Result<Vec<NodeResult>> {
1110 let mut handles = Vec::with_capacity(chunk.len());
1111 for prepared_task in chunk {
1112 let controller = controllers
1113 .get(&prepared_task.task.node_plan.controller_id)
1114 .ok_or_else(|| {
1115 DagMlError::RuntimeValidation(format!(
1116 "runtime controller `{}` is not registered",
1117 prepared_task.task.node_plan.controller_id
1118 ))
1119 })?;
1120 let worker_span = phase_span.clone();
1121 handles.push(thread_scope.spawn(move || {
1122 let _worker_span = worker_span.entered();
1123 let _node_span = crate::observability::node_span(
1124 prepared_task.task.run_id.as_str(),
1125 plan_id,
1126 prepared_task.task.phase.as_str(),
1127 prepared_task.task.node_plan.node_id.as_str(),
1128 prepared_task.task.node_plan.controller_id.as_str(),
1129 )
1130 .entered();
1131 let mut result = controller.invoke(&prepared_task.task)?;
1132 record_fit_influence_diagnostic(&prepared_task.task, &mut result);
1133 normalize_result_prediction_ports(
1134 plan,
1135 &prepared_task.task,
1136 &mut result,
1137 )?;
1138 result.validate_for_task(&prepared_task.task)?;
1139 Ok(result)
1140 }));
1141 }
1142 handles
1143 .into_iter()
1144 .map(|handle| {
1145 handle.join().map_err(|_| {
1146 DagMlError::RuntimeValidation(
1147 "parallel scheduler worker panicked".to_string(),
1148 )
1149 })?
1150 })
1151 .collect()
1152 })?;
1153
1154 for (prepared_task, mut result) in chunk.iter().zip(chunk_results) {
1155 apply_result_prediction_aggregation(
1156 plan,
1157 controllers,
1158 &prepared_task.task,
1159 &mut result,
1160 &resources,
1161 )?;
1162 attach_coordinator_input_lineage(
1163 &mut result,
1164 plan,
1165 &prepared_task.task.node_plan.node_id,
1166 &input_lineage,
1167 )?;
1168 if let Some(store) = resources.artifact_store.as_deref_mut() {
1169 if scope.phase == Phase::Refit {
1170 store.capture_refit_artifacts(&prepared_task.task, &result)?;
1171 }
1172 }
1173 for prediction in &result.predictions {
1174 ctx.prediction_store.append(prediction.clone())?;
1175 }
1176 for prediction in &result.aggregated_predictions {
1177 ctx.aggregated_prediction_store.append(prediction.clone())?;
1178 }
1179 apply_result_scoring(
1180 &result,
1181 &mut ctx.score_collector,
1182 &mut ctx.regression_target_records,
1183 )?;
1184 ctx.lineage.record(result.lineage.clone())?;
1185 let data_views = derive_output_data_views(plan, &prepared_task.task, &result)?;
1186 output_handles.insert(prepared_task.node_id.clone(), result.outputs.clone());
1187 output_data_views.insert(prepared_task.node_id.clone(), data_views);
1188 input_lineage.insert(
1189 prepared_task.node_id.clone(),
1190 result.lineage.record_id.clone(),
1191 );
1192 results.push(result);
1193 }
1194 }
1195
1196 for (node_id, reduction) in &merge_nodes {
1201 let node_plan = plan
1202 .node_plans
1203 .get(node_id)
1204 .expect("execution plan was validated");
1205 if let Some(mut result) =
1206 reassemble_branch_merge(plan, node_plan, ctx, &scope, *reduction)?
1207 {
1208 let task_node_plan = effective_node_plan_for_scope(node_plan, &scope)?;
1209 let task = NodeTask {
1210 inner_fold_set: None,
1211 run_id: ctx.run_id.clone(),
1212 node_plan: task_node_plan.clone(),
1213 phase: scope.phase,
1214 variant_id: scope.variant_id.clone(),
1215 variant: scope.variant.clone(),
1216 fold_id: scope.fold_id.clone(),
1217 branch_path: Vec::new(),
1218 input_handles: BTreeMap::new(),
1219 data_views: BTreeMap::new(),
1220 prediction_inputs: BTreeMap::new(),
1221 artifact_inputs: BTreeMap::new(),
1222 required_loss_attestations: NodeTask::required_loss_attestations_for(
1223 &task_node_plan,
1224 scope.phase,
1225 )?,
1226 fit_influence: FitInfluenceTask::default(),
1227 seed: None,
1228 };
1229 normalize_result_prediction_ports(plan, &task, &mut result)?;
1230 result.validate_for_task(&task)?;
1231 for prediction in &result.predictions {
1232 ctx.prediction_store.append(prediction.clone())?;
1233 }
1234 apply_result_scoring(
1235 &result,
1236 &mut ctx.score_collector,
1237 &mut ctx.regression_target_records,
1238 )?;
1239 ctx.lineage.record(result.lineage.clone())?;
1240 output_handles.insert(node_id.clone(), result.outputs.clone());
1241 input_lineage.insert(node_id.clone(), result.lineage.record_id.clone());
1242 results.push(result);
1243 }
1244 }
1245 }
1246
1247 Ok(results)
1248 }
1249}
1250
1251pub(crate) struct PreparedNodeTask {
1252 pub(crate) node_id: NodeId,
1253 pub(crate) task: NodeTask,
1254}
1255
1256pub(crate) fn attach_coordinator_input_lineage(
1257 result: &mut NodeResult,
1258 plan: &ExecutionPlan,
1259 node_id: &NodeId,
1260 upstream_lineage: &BTreeMap<NodeId, LineageId>,
1261) -> Result<()> {
1262 let inferred = inferred_input_lineage_for_node(plan, node_id, upstream_lineage);
1263 if result.lineage.input_lineage.is_empty() {
1264 result.lineage.input_lineage = inferred;
1265 return Ok(());
1266 }
1267
1268 let declared = result
1269 .lineage
1270 .input_lineage
1271 .iter()
1272 .cloned()
1273 .collect::<BTreeSet<_>>()
1274 .into_iter()
1275 .collect::<Vec<_>>();
1276 if declared != inferred {
1277 return Err(DagMlError::RuntimeValidation(format!(
1278 "lineage for node `{}` declared input lineage {:?}, expected {:?}",
1279 result.node_id, declared, inferred
1280 )));
1281 }
1282 result.lineage.input_lineage = declared;
1283 Ok(())
1284}
1285
1286pub(crate) fn inferred_input_lineage_for_node(
1287 plan: &ExecutionPlan,
1288 node_id: &NodeId,
1289 upstream_lineage: &BTreeMap<NodeId, LineageId>,
1290) -> Vec<LineageId> {
1291 plan.graph_plan
1292 .graph
1293 .edges
1294 .iter()
1295 .filter(|edge| &edge.target.node_id == node_id && edge.contract.propagates_lineage)
1296 .filter_map(|edge| upstream_lineage.get(&edge.source.node_id).cloned())
1297 .collect::<BTreeSet<_>>()
1298 .into_iter()
1299 .collect()
1300}
1301pub(crate) fn collect_input_handles(
1302 plan: &ExecutionPlan,
1303 node_plan: &NodePlan,
1304 output_handles: &BTreeMap<NodeId, BTreeMap<String, HandleRef>>,
1305 output_data_views: &BTreeMap<NodeId, BTreeMap<String, DataProviderViewSpec>>,
1306 resources: &PhaseScopeResources<'_>,
1307 ctx: &RunContext,
1308 scope: &PhaseScope,
1309) -> Result<CollectedInputs> {
1310 let mut inputs = BTreeMap::new();
1311 let mut data_views = BTreeMap::new();
1312 let mut prediction_inputs = BTreeMap::new();
1313 let training_oof_edges = incoming_training_oof_edges(plan, node_plan, scope)?;
1314 let masked_oof_source_ports = if scope.phase == Phase::Predict {
1320 incoming_oof_edges(plan, node_plan)?
1321 } else {
1322 training_oof_edges.clone()
1323 }
1324 .into_iter()
1325 .map(|edge| (edge.source.node_id.clone(), edge.source.port_name.clone()))
1326 .collect::<BTreeSet<_>>();
1327 let bound_data_inputs = node_plan
1328 .data_bindings
1329 .iter()
1330 .map(|binding| binding.input_name.clone())
1331 .collect::<BTreeSet<_>>();
1332 let declared_source_ports = plan
1336 .graph_plan
1337 .graph
1338 .edges
1339 .iter()
1340 .filter(|edge| edge.target.node_id == node_plan.node_id)
1341 .map(|edge| (edge.source.node_id.clone(), edge.source.port_name.clone()))
1342 .collect::<BTreeSet<_>>();
1343 for upstream in &node_plan.input_nodes {
1344 if let Some(handles) = output_handles.get(upstream) {
1345 for (port, handle) in handles {
1346 if !declared_source_ports.contains(&(upstream.clone(), port.clone())) {
1347 continue;
1348 }
1349 if masked_oof_source_ports.contains(&(upstream.clone(), port.clone())) {
1350 continue;
1351 }
1352 inputs.insert(format!("{upstream}.{port}"), handle.clone());
1353 }
1354 }
1355 }
1356 for edge in plan
1357 .graph_plan
1358 .graph
1359 .edges
1360 .iter()
1361 .filter(|edge| edge.target.node_id == node_plan.node_id)
1362 .filter(|edge| edge.contract.kind == PortKind::Data && !edge.contract.requires_oof)
1363 {
1364 if bound_data_inputs.contains(&edge.target.port_name) {
1365 continue;
1366 }
1367 let Some(handles) = output_handles.get(&edge.source.node_id) else {
1368 continue;
1369 };
1370 let Some(handle) = handles.get(&edge.source.port_name) else {
1371 continue;
1372 };
1373 let key = data_view_key(&edge.target.port_name);
1374 if inputs.insert(key.clone(), handle.clone()).is_some() {
1375 return Err(DagMlError::RuntimeValidation(format!(
1376 "node `{}` received duplicate data edge input `{key}`",
1377 node_plan.node_id
1378 )));
1379 }
1380 if let Some(source_views) = output_data_views.get(&edge.source.node_id) {
1381 if let Some(view) = source_views.get(&edge.source.port_name) {
1382 if data_views.insert(key.clone(), view.clone()).is_some() {
1383 return Err(DagMlError::RuntimeValidation(format!(
1384 "node `{}` received duplicate data edge view `{key}`",
1385 node_plan.node_id
1386 )));
1387 }
1388 }
1389 let source_validation_key = validation_data_view_key(&edge.source.port_name);
1390 if let Some(view) = source_views.get(&source_validation_key) {
1391 let validation_key = format!("{key}:validation");
1392 if data_views
1393 .insert(validation_key.clone(), view.clone())
1394 .is_some()
1395 {
1396 return Err(DagMlError::RuntimeValidation(format!(
1397 "node `{}` received duplicate data edge validation view `{validation_key}`",
1398 node_plan.node_id
1399 )));
1400 }
1401 }
1402 }
1403 }
1404 for edge in training_oof_edges {
1405 let key = format!("{}.{}", edge.source.node_id, edge.source.port_name);
1406 let Some(input) = collect_oof_prediction_input(plan, edge, ctx, scope, resources)? else {
1407 return Ok(CollectedInputs {
1408 handles: BTreeMap::new(),
1409 data_views: BTreeMap::new(),
1410 prediction_inputs: BTreeMap::new(),
1411 skip_node: true,
1412 });
1413 };
1414 if inputs.insert(key.clone(), input.handle).is_some() {
1415 return Err(DagMlError::RuntimeValidation(format!(
1416 "node `{}` received duplicate OOF prediction input `{key}`",
1417 node_plan.node_id
1418 )));
1419 }
1420 if prediction_inputs.insert(key.clone(), input.spec).is_some() {
1421 return Err(DagMlError::RuntimeValidation(format!(
1422 "node `{}` received duplicate OOF prediction spec `{key}`",
1423 node_plan.node_id
1424 )));
1425 }
1426 }
1427 if matches!(scope.phase, Phase::Refit | Phase::Predict) {
1434 let off_fold_suffix = scope.phase.as_str().to_ascii_lowercase();
1435 for edge in incoming_oof_edges(plan, node_plan)? {
1436 let Some(input) = collect_off_fold_prediction_input(plan, edge, ctx, scope)? else {
1437 continue;
1438 };
1439 let key = format!(
1440 "{}.{}:{off_fold_suffix}",
1441 edge.source.node_id, edge.source.port_name
1442 );
1443 if inputs.insert(key.clone(), input.handle).is_some() {
1444 return Err(DagMlError::RuntimeValidation(format!(
1445 "node `{}` received duplicate off-fold prediction input `{key}`",
1446 node_plan.node_id
1447 )));
1448 }
1449 if prediction_inputs.insert(key.clone(), input.spec).is_some() {
1450 return Err(DagMlError::RuntimeValidation(format!(
1451 "node `{}` received duplicate off-fold prediction spec `{key}`",
1452 node_plan.node_id
1453 )));
1454 }
1455 }
1456 }
1457 if !node_plan.data_bindings.is_empty() && resources.data_provider.is_none() {
1458 return Err(DagMlError::RuntimeValidation(format!(
1459 "node `{}` requires {} data binding(s) but no runtime data provider is registered",
1460 node_plan.node_id,
1461 node_plan.data_bindings.len()
1462 )));
1463 }
1464 if let Some(data_provider) = resources.data_provider {
1465 let excluded_samples = coordinator_relations_for_node(node_plan, resources)?
1469 .map(|relations| relations.excluded_sample_ids())
1470 .unwrap_or_default();
1471 for binding in &node_plan.data_bindings {
1472 let materialized = data_provider.materialize(&DataMaterializationRequest {
1473 run_id: ctx.run_id.clone(),
1474 node_id: node_plan.node_id.clone(),
1475 input_name: binding.input_name.clone(),
1476 phase: scope.phase,
1477 variant_id: scope.variant_id.clone(),
1478 fold_id: scope.fold_id.clone(),
1479 binding: binding.clone(),
1480 })?;
1481 let branch_view_for_node = branch_view_from_node_metadata(plan, &node_plan.node_id)?;
1482 let view = data_view_for_scope(
1483 binding,
1484 plan.fold_set.as_ref(),
1485 scope,
1486 branch_view_for_node.as_ref(),
1487 &excluded_samples,
1488 )?;
1489 let key = data_view_key(&binding.input_name);
1490 let view_handle = make_data_view_handle(
1491 data_provider,
1492 ctx,
1493 node_plan,
1494 scope,
1495 binding,
1496 &materialized,
1497 &view,
1498 )?;
1499 if data_views.insert(key.clone(), view).is_some() {
1500 return Err(DagMlError::RuntimeValidation(format!(
1501 "node `{}` received duplicate data view `{key}`",
1502 node_plan.node_id
1503 )));
1504 }
1505 if inputs.insert(key.clone(), view_handle).is_some() {
1506 return Err(DagMlError::RuntimeValidation(format!(
1507 "node `{}` received duplicate data input `{key}`",
1508 node_plan.node_id
1509 )));
1510 }
1511
1512 if let Some(validation_view) = validation_data_view_for_scope(
1513 binding,
1514 plan.fold_set.as_ref(),
1515 scope,
1516 branch_view_for_node.as_ref(),
1517 &excluded_samples,
1518 )? {
1519 let validation_key = format!("{key}:validation");
1520 let validation_handle = make_data_view_handle(
1521 data_provider,
1522 ctx,
1523 node_plan,
1524 scope,
1525 binding,
1526 &materialized,
1527 &validation_view,
1528 )?;
1529 if data_views
1530 .insert(validation_key.clone(), validation_view)
1531 .is_some()
1532 {
1533 return Err(DagMlError::RuntimeValidation(format!(
1534 "node `{}` received duplicate validation data view `{validation_key}`",
1535 node_plan.node_id
1536 )));
1537 }
1538 if inputs
1539 .insert(validation_key.clone(), validation_handle)
1540 .is_some()
1541 {
1542 return Err(DagMlError::RuntimeValidation(format!(
1543 "node `{}` received duplicate validation data input `{validation_key}`",
1544 node_plan.node_id
1545 )));
1546 }
1547 }
1548 }
1549 }
1550 Ok(CollectedInputs {
1551 handles: inputs,
1552 data_views,
1553 prediction_inputs,
1554 skip_node: false,
1555 })
1556}
1557pub(crate) fn preload_replay_prediction_cache_store(
1558 bundle: &ExecutionBundle,
1559 prediction_cache_store: Option<&dyn RuntimePredictionCacheStore>,
1560 ctx: &mut RunContext,
1561) -> Result<()> {
1562 if bundle.prediction_requirements.is_empty() {
1563 return Ok(());
1564 }
1565 let store = prediction_cache_store.ok_or_else(|| {
1566 DagMlError::RuntimeValidation(format!(
1567 "bundle `{}` cannot preload OOF prediction caches without a prediction cache store",
1568 bundle.bundle_id
1569 ))
1570 })?;
1571 if !ctx.prediction_store.blocks().is_empty() {
1572 return Err(DagMlError::RuntimeValidation(format!(
1573 "bundle `{}` cannot preload OOF prediction caches into a non-empty prediction store",
1574 bundle.bundle_id
1575 )));
1576 }
1577 let contracts = replay_prediction_cache_contracts(bundle)?;
1578 for contract in contracts.values() {
1579 if contract.requirement.prediction_level == PredictionLevel::Sample {
1580 let blocks = store.load_blocks(&contract.cache.requirement_key)?;
1581 if blocks.iter().any(|block| {
1582 block.producer_node != contract.requirement.producer_node
1583 || block.partition != contract.requirement.partition
1584 }) {
1585 return Err(DagMlError::RuntimeValidation(format!(
1586 "prediction cache store returned blocks outside requirement `{}`",
1587 contract.cache.requirement_key
1588 )));
1589 }
1590 let mut payload = build_prediction_cache_payload(&contract.requirement, &blocks)?;
1591 payload.cache_namespace_fingerprints =
1592 contract.cache.cache_namespace_fingerprints.clone();
1593 validate_prediction_cache_payload_matches_record(&payload, &contract.cache)?;
1594 for block in &payload.blocks {
1595 ctx.prediction_store.append(block.clone())?;
1596 }
1597 } else {
1598 let blocks = store.load_aggregated_blocks(&contract.cache.requirement_key)?;
1599 if blocks.iter().any(|block| {
1600 block.producer_node != contract.requirement.producer_node
1601 || block.partition != contract.requirement.partition
1602 || block.level != contract.requirement.prediction_level
1603 }) {
1604 return Err(DagMlError::RuntimeValidation(format!(
1605 "prediction cache store returned aggregated blocks outside requirement `{}`",
1606 contract.cache.requirement_key
1607 )));
1608 }
1609 let mut payload =
1610 build_aggregated_prediction_cache_payload(&contract.requirement, &blocks)?;
1611 payload.cache_namespace_fingerprints =
1612 contract.cache.cache_namespace_fingerprints.clone();
1613 validate_prediction_cache_payload_matches_record(&payload, &contract.cache)?;
1614 }
1615 }
1616 Ok(())
1617}
1618
1619pub(crate) fn replay_prediction_cache_contracts(
1620 bundle: &ExecutionBundle,
1621) -> Result<BTreeMap<String, ReplayPredictionCacheContract>> {
1622 bundle.validate()?;
1623 let requirements = bundle
1624 .prediction_requirements
1625 .iter()
1626 .map(|requirement| (requirement.key(), requirement))
1627 .collect::<BTreeMap<_, _>>();
1628 let mut contracts = BTreeMap::new();
1629 for cache in &bundle.prediction_caches {
1630 let requirement = requirements.get(&cache.requirement_key).ok_or_else(|| {
1631 DagMlError::RuntimeValidation(format!(
1632 "prediction cache `{}` references unknown prediction requirement `{}`",
1633 cache.cache_id, cache.requirement_key
1634 ))
1635 })?;
1636 contracts.insert(
1637 cache.requirement_key.clone(),
1638 ReplayPredictionCacheContract {
1639 requirement: (*requirement).clone(),
1640 cache: cache.clone(),
1641 },
1642 );
1643 }
1644 Ok(contracts)
1645}
1646
1647pub(crate) fn materialize_replay_artifact_handles(
1648 plan: &ExecutionPlan,
1649 bundle: &ExecutionBundle,
1650 replay_request: &ReplayPhaseRequest,
1651 artifact_store: &dyn RuntimeArtifactStore,
1652 ctx: &RunContext,
1653) -> Result<MaterializedReplayArtifacts> {
1654 let mut handles = BTreeMap::<NodeId, BTreeMap<String, HandleRef>>::new();
1655 let mut inputs = BTreeMap::<NodeId, BTreeMap<String, ArtifactInputSpec>>::new();
1656 for artifact in &bundle.refit_artifacts {
1657 artifact.validate()?;
1658 let node_plan = plan.node_plans.get(&artifact.node_id).ok_or_else(|| {
1659 DagMlError::RuntimeValidation(format!(
1660 "bundle `{}` artifact references unknown node `{}`",
1661 bundle.bundle_id, artifact.node_id
1662 ))
1663 })?;
1664 if !node_plan.supported_phases.contains(&replay_request.phase) {
1665 return Err(DagMlError::RuntimeValidation(format!(
1666 "bundle `{}` artifact node `{}` does not support replay phase {:?}",
1667 bundle.bundle_id, artifact.node_id, replay_request.phase
1668 )));
1669 }
1670 let handle = artifact_store.materialize(&ArtifactMaterializationRequest {
1671 run_id: ctx.run_id.clone(),
1672 bundle_id: bundle.bundle_id.clone(),
1673 node_id: artifact.node_id.clone(),
1674 phase: replay_request.phase,
1675 variant_id: bundle.selected_variant_id.clone(),
1676 controller_id: artifact.controller_id.clone(),
1677 artifact: artifact.artifact.clone(),
1678 params_fingerprint: artifact.params_fingerprint.clone(),
1679 training_loss_fingerprint: artifact.training_loss_fingerprint.clone(),
1680 })?;
1681 if !matches!(handle.kind, HandleKind::Model | HandleKind::Artifact) {
1682 return Err(DagMlError::RuntimeValidation(format!(
1683 "artifact `{}` materialized as unsupported handle kind {:?}",
1684 artifact.artifact.id, handle.kind
1685 )));
1686 }
1687 if handle.owner_controller != artifact.controller_id {
1688 return Err(DagMlError::RuntimeValidation(format!(
1689 "artifact `{}` handle owner `{}` does not match controller `{}`",
1690 artifact.artifact.id, handle.owner_controller, artifact.controller_id
1691 )));
1692 }
1693 let key = refit_artifact_input_key(&artifact.artifact.id);
1694 if handles
1695 .entry(artifact.node_id.clone())
1696 .or_default()
1697 .insert(key.clone(), handle)
1698 .is_some()
1699 {
1700 return Err(DagMlError::RuntimeValidation(format!(
1701 "duplicate replay artifact input `{key}` for node `{}`",
1702 artifact.node_id
1703 )));
1704 }
1705 if inputs
1706 .entry(artifact.node_id.clone())
1707 .or_default()
1708 .insert(key.clone(), ArtifactInputSpec::from_refit_record(artifact)?)
1709 .is_some()
1710 {
1711 return Err(DagMlError::RuntimeValidation(format!(
1712 "duplicate replay artifact metadata `{key}` for node `{}`",
1713 artifact.node_id
1714 )));
1715 }
1716 }
1717 Ok(MaterializedReplayArtifacts { handles, inputs })
1718}
1719
1720pub(crate) fn derive_task_seed(
1721 root_seed: Option<u64>,
1722 variant_id: Option<&VariantId>,
1723 fold_id: Option<&FoldId>,
1724 node_plan: &NodePlan,
1725 phase: Phase,
1726) -> Option<u64> {
1727 root_seed.map(|root| {
1728 let mut context = SeedContext::root(root);
1729 if let Some(variant_id) = variant_id {
1730 context = context.child(format!("variant:{variant_id}"));
1731 }
1732 if let Some(fold_id) = fold_id {
1733 context = context.child(format!("fold:{fold_id}"));
1734 }
1735 context
1736 .child(format!("node:{}", node_plan.node_id))
1737 .child(format!("phase:{phase:?}"))
1738 .derive_u64("task")
1739 })
1740}