1use std::any::Any;
2use std::collections::{BTreeMap, BTreeSet};
3use std::fmt;
4use std::panic::{catch_unwind, AssertUnwindSafe};
5use std::sync::{Arc, OnceLock};
6use std::time::{Duration, Instant};
7
8use super::{
9 defer_device_cleanup, deferred_device_cleanup_status, maintain_deferred_device_cleanups,
10 new_deferred_device_cleanup_domain, BufferUsage, DeferredDeviceCleanupDisposition,
11 DeferredDeviceCleanupDomainId, DeferredDeviceCleanupMaintenanceReceipt,
12 DeferredDeviceCleanupStatus, DeferredDeviceCleanupTask, DeviceRuntime, ElementType,
13 ExecutionPlan, FailureDomain, FailureEnvelope, PlanRuntimeHandoffError, PlanRuntimeResources,
14 ResourceId, ResourceTransaction, ResourceTransactionDriver, TransactionCommitted, VNextError,
15 MAX_DEFERRED_DEVICE_CLEANUP_MAINTENANCE_TASKS,
16};
17use super::{
18 DeviceCommandBatch, DeviceTerminal, HostTransferLayout, PreparedModelFamily,
19 StaticWeightTransformDestination, StaticWeightTransformPlan, StaticWeightTransformRequest,
20 WeightComponentPayload, WeightComponentSegments, WeightComponentSource, WeightComponentSpec,
21 WeightId,
22};
23
24static STATIC_INITIALIZATION_CLEANUP_DOMAIN: OnceLock<DeferredDeviceCleanupDomainId> =
25 OnceLock::new();
26
27fn static_initialization_cleanup_domain() -> DeferredDeviceCleanupDomainId {
28 *STATIC_INITIALIZATION_CLEANUP_DOMAIN.get_or_init(new_deferred_device_cleanup_domain)
29}
30
31pub fn static_initialization_cleanup_status() -> DeferredDeviceCleanupStatus {
34 deferred_device_cleanup_status(static_initialization_cleanup_domain())
35}
36
37pub fn maintain_static_initialization_cleanups(
40 maximum_tasks: usize,
41) -> Result<DeferredDeviceCleanupMaintenanceReceipt, VNextError> {
42 if maximum_tasks == 0 || maximum_tasks > MAX_DEFERRED_DEVICE_CLEANUP_MAINTENANCE_TASKS {
43 return Err(VNextError::InvalidExecutionPlan {
44 reason: format!(
45 "static initialization cleanup maintenance size must be in 1..={MAX_DEFERRED_DEVICE_CLEANUP_MAINTENANCE_TASKS}"
46 ),
47 });
48 }
49 Ok(maintain_deferred_device_cleanups(
50 static_initialization_cleanup_domain(),
51 maximum_tasks,
52 ))
53}
54
55#[derive(Debug, Clone, Copy, PartialEq, Eq)]
59pub struct StaticInitializationPolicy {
60 maximum_staging_bytes: u64,
61 maximum_commands_per_batch: usize,
62}
63
64impl StaticInitializationPolicy {
65 pub fn new(
66 maximum_staging_bytes: u64,
67 maximum_commands_per_batch: usize,
68 ) -> Result<Self, VNextError> {
69 if maximum_staging_bytes == 0 || maximum_commands_per_batch == 0 {
70 return Err(VNextError::InvalidExecutionPlan {
71 reason: "static initialization requires non-zero staging and command budgets"
72 .to_owned(),
73 });
74 }
75 Ok(Self {
76 maximum_staging_bytes,
77 maximum_commands_per_batch,
78 })
79 }
80
81 pub const fn maximum_staging_bytes(self) -> u64 {
82 self.maximum_staging_bytes
83 }
84
85 pub const fn maximum_commands_per_batch(self) -> usize {
86 self.maximum_commands_per_batch
87 }
88}
89
90#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
91pub struct StaticInitializationReceipt {
92 initialized_resource_count: usize,
93 uploaded_component_count: usize,
94 uploaded_bytes: u64,
95 imported_component_count: usize,
96 imported_bytes: u64,
97 transformed_component_count: usize,
98 transformed_bytes: u64,
99 transform_command_count: usize,
100 upload_command_count: usize,
101 submission_batch_count: usize,
102 total_duration_us: u64,
103 setup_duration_us: u64,
104 source_materialization_duration_us: u64,
105 device_encode_duration_us: u64,
106 device_import_duration_us: u64,
107 device_transform_encode_duration_us: u64,
108 submission_wait_duration_us: u64,
109 import_seal_duration_us: u64,
110 slowest_component_id: Option<WeightId>,
111 slowest_component_materialization_duration_us: u64,
112 source_files: BTreeSet<String>,
113}
114
115impl StaticInitializationReceipt {
116 pub const fn initialized_resource_count(&self) -> usize {
117 self.initialized_resource_count
118 }
119
120 pub const fn uploaded_component_count(&self) -> usize {
121 self.uploaded_component_count
122 }
123
124 pub const fn uploaded_bytes(&self) -> u64 {
125 self.uploaded_bytes
126 }
127
128 pub const fn imported_component_count(&self) -> usize {
129 self.imported_component_count
130 }
131
132 pub const fn imported_bytes(&self) -> u64 {
133 self.imported_bytes
134 }
135
136 pub const fn transformed_component_count(&self) -> usize {
137 self.transformed_component_count
138 }
139
140 pub const fn transformed_bytes(&self) -> u64 {
141 self.transformed_bytes
142 }
143
144 pub const fn transform_command_count(&self) -> usize {
145 self.transform_command_count
146 }
147
148 pub const fn upload_command_count(&self) -> usize {
149 self.upload_command_count
150 }
151
152 pub const fn submission_batch_count(&self) -> usize {
153 self.submission_batch_count
154 }
155
156 pub const fn total_duration_us(&self) -> u64 {
157 self.total_duration_us
158 }
159
160 pub const fn setup_duration_us(&self) -> u64 {
161 self.setup_duration_us
162 }
163
164 pub const fn source_materialization_duration_us(&self) -> u64 {
165 self.source_materialization_duration_us
166 }
167
168 pub const fn device_encode_duration_us(&self) -> u64 {
169 self.device_encode_duration_us
170 }
171
172 pub const fn device_import_duration_us(&self) -> u64 {
173 self.device_import_duration_us
174 }
175
176 pub const fn device_transform_encode_duration_us(&self) -> u64 {
177 self.device_transform_encode_duration_us
178 }
179
180 pub const fn submission_wait_duration_us(&self) -> u64 {
181 self.submission_wait_duration_us
182 }
183
184 pub const fn import_seal_duration_us(&self) -> u64 {
185 self.import_seal_duration_us
186 }
187
188 pub fn slowest_component_id(&self) -> Option<&WeightId> {
189 self.slowest_component_id.as_ref()
190 }
191
192 pub const fn slowest_component_materialization_duration_us(&self) -> u64 {
193 self.slowest_component_materialization_duration_us
194 }
195
196 pub fn source_files(&self) -> &BTreeSet<String> {
197 &self.source_files
198 }
199}
200
201#[must_use = "initialized static resources must be handed to the plan runtime"]
205pub struct InitializedResourceTransaction<D>
206where
207 D: ResourceTransactionDriver,
208{
209 transaction: ResourceTransaction<D, TransactionCommitted>,
210 receipt: StaticInitializationReceipt,
211}
212
213impl<D> InitializedResourceTransaction<D>
214where
215 D: ResourceTransactionDriver,
216{
217 pub fn receipt(&self) -> &StaticInitializationReceipt {
218 &self.receipt
219 }
220
221 pub fn into_plan_runtime(
222 self,
223 ) -> Result<Arc<PlanRuntimeResources<D::Runtime>>, PlanRuntimeHandoffError<D>>
224 where
225 D: 'static,
226 {
227 self.transaction.into_plan_runtime()
228 }
229}
230
231struct StaticInitializationRecovery<R>
232where
233 R: DeviceRuntime,
234{
235 stream: R::Stream,
236 fence: Option<R::Fence>,
237}
238
239#[must_use = "static initialization failure retains transaction and possibly in-flight ownership"]
244pub struct StaticInitializationFailure<D>
245where
246 D: ResourceTransactionDriver + 'static,
247{
248 transaction: Option<ResourceTransaction<D, TransactionCommitted>>,
249 failure: FailureEnvelope,
250 recovery: Option<StaticInitializationRecovery<D::Runtime>>,
251}
252
253impl<D> fmt::Debug for StaticInitializationFailure<D>
254where
255 D: ResourceTransactionDriver + 'static,
256{
257 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
258 formatter
259 .debug_struct("StaticInitializationFailure")
260 .field("failure", &self.failure)
261 .field("indeterminate", &self.recovery.is_some())
262 .finish_non_exhaustive()
263 }
264}
265
266impl<D> fmt::Display for StaticInitializationFailure<D>
267where
268 D: ResourceTransactionDriver + 'static,
269{
270 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
271 write!(
272 formatter,
273 "static initialization failed: {}",
274 self.failure.message()
275 )
276 }
277}
278
279impl<D> std::error::Error for StaticInitializationFailure<D> where
280 D: ResourceTransactionDriver + 'static
281{
282}
283
284impl<D> StaticInitializationFailure<D>
285where
286 D: ResourceTransactionDriver + 'static,
287{
288 fn new(
289 transaction: ResourceTransaction<D, TransactionCommitted>,
290 step: InitializationStepFailure<D::Runtime>,
291 ) -> Self {
292 match step {
293 InitializationStepFailure::Quiescent(failure) => Self {
294 transaction: Some(transaction),
295 failure,
296 recovery: None,
297 },
298 InitializationStepFailure::Indeterminate { failure, recovery } => Self {
299 transaction: Some(transaction),
300 failure,
301 recovery: Some(recovery),
302 },
303 }
304 }
305
306 pub fn failure(&self) -> &FailureEnvelope {
307 &self.failure
308 }
309
310 pub const fn is_indeterminate(&self) -> bool {
311 self.recovery.is_some()
312 }
313
314 pub fn into_transaction(
315 mut self,
316 ) -> Result<ResourceTransaction<D, TransactionCommitted>, Self> {
317 if self.recovery.is_some() {
318 return Err(self);
319 }
320 Ok(self
321 .transaction
322 .take()
323 .expect("static initialization failure owns its transaction"))
324 }
325
326 pub fn recover(mut self) -> Result<ResourceTransaction<D, TransactionCommitted>, Self> {
329 let Some(mut recovery) = self.recovery.take() else {
330 return Ok(self
331 .transaction
332 .take()
333 .expect("static initialization failure owns its transaction"));
334 };
335 let runtime = Arc::clone(
336 self.transaction
337 .as_ref()
338 .expect("static initialization failure owns its transaction")
339 .lease()
340 .runtime(),
341 );
342 let synchronized = catch_unwind(AssertUnwindSafe(|| {
343 runtime.synchronize(&mut recovery.stream)
344 }));
345 match synchronized {
346 Ok(Ok(())) => {
347 drop(recovery.fence.take());
348 Ok(self
349 .transaction
350 .take()
351 .expect("static initialization failure owns its transaction"))
352 }
353 Ok(Err(error)) => {
354 self.failure = device_failure(&runtime, &error, "static_recovery");
355 self.recovery = Some(recovery);
356 Err(self)
357 }
358 Err(payload) => {
359 self.failure = portable_failure(
360 FailureDomain::Device,
361 "static_recovery_panic",
362 panic_message(payload),
363 false,
364 );
365 self.recovery = Some(recovery);
366 Err(self)
367 }
368 }
369 }
370}
371
372impl<D> Drop for StaticInitializationFailure<D>
373where
374 D: ResourceTransactionDriver + 'static,
375{
376 fn drop(&mut self) {
377 if let Some(recovery) = self.recovery.take() {
378 let transaction = self
379 .transaction
380 .take()
381 .expect("indeterminate static initialization owns its transaction");
382 defer_device_cleanup(
383 static_initialization_cleanup_domain(),
384 DeferredStaticInitializationCleanup {
385 transaction: Some(transaction),
386 recovery: Some(recovery),
387 },
388 );
389 }
390 }
391}
392
393struct DeferredStaticInitializationCleanup<D>
394where
395 D: ResourceTransactionDriver + 'static,
396{
397 transaction: Option<ResourceTransaction<D, TransactionCommitted>>,
398 recovery: Option<StaticInitializationRecovery<D::Runtime>>,
399}
400
401impl<D> DeferredDeviceCleanupTask for DeferredStaticInitializationCleanup<D>
402where
403 D: ResourceTransactionDriver + 'static,
404{
405 fn try_cleanup(&mut self) -> DeferredDeviceCleanupDisposition {
406 let transaction = self
407 .transaction
408 .as_ref()
409 .expect("deferred static initialization owns its transaction");
410 let recovery = self
411 .recovery
412 .as_mut()
413 .expect("deferred static initialization owns its recovery stream");
414 let runtime = Arc::clone(transaction.lease().runtime());
415 let synchronized = catch_unwind(AssertUnwindSafe(|| {
416 runtime.synchronize(&mut recovery.stream)
417 }));
418 if !matches!(synchronized, Ok(Ok(()))) {
419 return DeferredDeviceCleanupDisposition::Retryable;
420 }
421 let mut recovery = self
422 .recovery
423 .take()
424 .expect("successful recovery retains its stream and fence");
425 drop(recovery.fence.take());
426 drop(recovery);
427 drop(
428 self.transaction
429 .take()
430 .expect("successful recovery retains its transaction"),
431 );
432 DeferredDeviceCleanupDisposition::Completed
433 }
434}
435
436#[derive(Debug, Clone, PartialEq, Eq)]
437struct WeightPlacement {
438 component_id: WeightId,
439 resource_id: ResourceId,
440 offset_bytes: u64,
441 length_bytes: u64,
442 element_type: ElementType,
443}
444
445enum InitializationStepFailure<R>
446where
447 R: DeviceRuntime,
448{
449 Quiescent(FailureEnvelope),
450 Indeterminate {
451 failure: FailureEnvelope,
452 recovery: StaticInitializationRecovery<R>,
453 },
454}
455
456impl<D> ResourceTransaction<D, TransactionCommitted>
457where
458 D: ResourceTransactionDriver + 'static,
459{
460 pub fn initialize_static(
461 self,
462 family: &PreparedModelFamily,
463 plan: &ExecutionPlan,
464 source: &dyn WeightComponentSource,
465 policy: StaticInitializationPolicy,
466 ) -> Result<InitializedResourceTransaction<D>, StaticInitializationFailure<D>> {
467 match initialize_static_inner(&self, family, plan, source, policy) {
468 Ok(receipt) => Ok(InitializedResourceTransaction {
469 transaction: self,
470 receipt,
471 }),
472 Err(step) => Err(StaticInitializationFailure::new(self, step)),
473 }
474 }
475}
476
477fn initialize_static_inner<D>(
478 transaction: &ResourceTransaction<D, TransactionCommitted>,
479 family: &PreparedModelFamily,
480 plan: &ExecutionPlan,
481 source: &dyn WeightComponentSource,
482 policy: StaticInitializationPolicy,
483) -> Result<StaticInitializationReceipt, InitializationStepFailure<D::Runtime>>
484where
485 D: ResourceTransactionDriver,
486{
487 let initialization_started = Instant::now();
488 let setup_started = Instant::now();
489 preflight_transaction(transaction, family, plan).map_err(contract_failure)?;
490 let placements = weight_placements(family, plan).map_err(contract_failure)?;
491 let execution_weight_schema = plan.payload().execution_weights().schema();
492 let runtime = Arc::clone(transaction.lease().runtime());
493 let has_required_weight_transforms = !plan
494 .payload()
495 .execution_weights()
496 .static_weight_transforms()
497 .is_empty();
498 let mut weight_import = if placements.is_empty() || has_required_weight_transforms {
499 None
500 } else {
501 match runtime.begin_static_weight_import() {
502 None => None,
503 Some(Ok(import)) => Some(import),
504 Some(Err(error)) => {
505 return Err(InitializationStepFailure::Quiescent(device_failure(
506 &runtime,
507 &error,
508 "static_weight_import_begin",
509 )))
510 }
511 }
512 };
513 let created_stream = runtime.create_stream().map_err(|error| {
514 InitializationStepFailure::Quiescent(device_failure(
515 &runtime,
516 &error,
517 "static_stream_create",
518 ))
519 })?;
520 let mut stream = Some(created_stream);
521 let setup_duration = setup_started.elapsed();
522 let mut pending = Vec::<<D::Runtime as DeviceRuntime>::Command>::new();
523 let mut pending_staging_bytes = 0_u64;
524 let mut submission_batch_count = 0_usize;
525 let mut upload_command_count = 0_usize;
526 let mut uploaded_component_count = 0_usize;
527 let mut uploaded_bytes = 0_u64;
528 let mut imported_component_count = 0_usize;
529 let mut imported_bytes = 0_u64;
530 let mut transformed_component_count = 0_usize;
531 let mut transformed_bytes = 0_u64;
532 let mut transform_command_count = 0_usize;
533 let mut source_materialization_duration = Duration::ZERO;
534 let mut device_encode_duration = Duration::ZERO;
535 let mut device_import_duration = Duration::ZERO;
536 let mut device_transform_encode_duration = Duration::ZERO;
537 let mut submission_wait_duration = Duration::ZERO;
538 let mut import_seal_duration = Duration::ZERO;
539 let mut slowest_component_id = None;
540 let mut slowest_component_materialization_duration = Duration::ZERO;
541 let mut source_files = BTreeSet::new();
542
543 for allocation in plan.payload().memory().static_allocations() {
544 if allocation.usage() == BufferUsage::Weights && weight_import.is_some() {
545 continue;
546 }
547 let encode_started = Instant::now();
548 let command = with_static_buffer(transaction, allocation.resource_id(), |buffer| {
549 runtime.encode_zero(buffer, 0, allocation.size_bytes())
550 })
551 .map_err(|error| runtime_or_contract_failure(&runtime, error, "static_zero_encode"))?;
552 device_encode_duration += encode_started.elapsed();
553 pending.push(command);
554 if pending.len() == policy.maximum_commands_per_batch() {
555 submission_wait_duration += submit_pending(
556 &runtime,
557 &mut stream,
558 &mut pending,
559 &mut pending_staging_bytes,
560 )?;
561 submission_batch_count += 1;
562 }
563 }
564
565 let mut materialization_groups = Vec::<Vec<&WeightComponentSpec>>::new();
566 let mut materialization_group_indices = BTreeMap::<Vec<WeightId>, usize>::new();
567 for component in &execution_weight_schema.components {
568 if !placements.contains_key(&component.id) {
569 continue;
570 }
571 let source_ids = plan
572 .payload()
573 .execution_weights()
574 .component_sources()
575 .get(&component.id)
576 .ok_or_else(|| {
577 contract_failure(VNextError::InvalidExecutionPlan {
578 reason: format!(
579 "execution weight component `{}` has no source mapping",
580 component.id
581 ),
582 })
583 })?;
584 if let Some(group_index) = materialization_group_indices.get(source_ids) {
585 materialization_groups[*group_index].push(component);
586 } else {
587 let group_index = materialization_groups.len();
588 materialization_group_indices.insert(source_ids.clone(), group_index);
589 materialization_groups.push(vec![component]);
590 }
591 }
592
593 for components in materialization_groups {
594 if let Some(transform) = plan
595 .static_weight_transform_for_components(&components)
596 .map_err(contract_failure)?
597 {
598 let materialization_started = Instant::now();
599 let sources =
600 prepare_transform_sources(family, source, transform).map_err(contract_failure)?;
601 let materialization_duration = materialization_started.elapsed();
602 source_materialization_duration += materialization_duration;
603 if slowest_component_id.is_none()
604 || materialization_duration > slowest_component_materialization_duration
605 {
606 slowest_component_materialization_duration = materialization_duration;
607 slowest_component_id = Some(components[0].id.clone());
608 }
609 for source_segments in &sources {
610 source_files.extend(source_segments.source_files().iter().cloned());
611 }
612 let scratch_resource_id = plan
613 .payload()
614 .execution_weights()
615 .static_weight_transform_scratch_resource_id()
616 .map_err(contract_failure)?
617 .ok_or_else(|| {
618 contract_failure(VNextError::InvalidExecutionPlan {
619 reason: "required static weight transform has no admitted scratch resource"
620 .to_owned(),
621 })
622 })?;
623 let encode_started = Instant::now();
624 let command = encode_required_weight_transform(
625 transaction,
626 &runtime,
627 transform,
628 &sources,
629 &components,
630 &placements,
631 &scratch_resource_id,
632 )
633 .map_err(|error| {
634 runtime_or_contract_failure(&runtime, error, "static_weight_transform_encode")
635 })?;
636 device_transform_encode_duration += encode_started.elapsed();
637 pending.push(command);
638 transform_command_count += 1;
639 transformed_component_count += components.len();
640 transformed_bytes = components
641 .iter()
642 .try_fold(transformed_bytes, |total, component| {
643 total.checked_add(
644 placements
645 .get(&component.id)
646 .expect("transform components have selected placements")
647 .length_bytes,
648 )
649 })
650 .ok_or_else(|| {
651 contract_failure(VNextError::InvalidExecutionPlan {
652 reason: "static transformed bytes overflow u64".to_owned(),
653 })
654 })?;
655 if pending.len() == policy.maximum_commands_per_batch() {
656 submission_wait_duration += submit_pending(
657 &runtime,
658 &mut stream,
659 &mut pending,
660 &mut pending_staging_bytes,
661 )?;
662 submission_batch_count += 1;
663 }
664 continue;
665 }
666 let materialization_started = Instant::now();
670 let uploads = prepare_uploads(family, plan, source, &components, &placements)
671 .map_err(contract_failure)?;
672 let materialization_duration = materialization_started.elapsed();
673 source_materialization_duration += materialization_duration;
674 if slowest_component_id.is_none()
675 || materialization_duration > slowest_component_materialization_duration
676 {
677 slowest_component_materialization_duration = materialization_duration;
678 slowest_component_id = Some(components[0].id.clone());
679 }
680 for (component, upload) in components.into_iter().zip(uploads) {
681 let placement = placements
682 .get(&component.id)
683 .expect("materialization groups contain only placed components");
684 source_files.extend(upload.source_files().iter().cloned());
685 if let Some(import) = weight_import.as_mut() {
686 let import_started = Instant::now();
687 with_static_buffer(transaction, &placement.resource_id, |buffer| {
688 import.import_component(&upload, buffer, placement.offset_bytes)
689 })
690 .map_err(|error| {
691 runtime_or_contract_failure(&runtime, error, "static_weight_component_import")
692 })?;
693 device_import_duration += import_started.elapsed();
694 imported_component_count += 1;
695 imported_bytes = imported_bytes
696 .checked_add(placement.length_bytes)
697 .ok_or_else(|| {
698 contract_failure(VNextError::InvalidExecutionPlan {
699 reason: "static initialization imported bytes overflow u64".to_owned(),
700 })
701 })?;
702 continue;
703 }
704 let element_bytes = upload.element_type().size_bytes();
705 let maximum_chunk_bytes =
706 policy.maximum_staging_bytes() - policy.maximum_staging_bytes() % element_bytes;
707 if maximum_chunk_bytes == 0 {
708 return Err(contract_failure(VNextError::InvalidExecutionPlan {
709 reason: format!(
710 "static staging budget cannot hold one {:?} element",
711 upload.element_type()
712 ),
713 }));
714 }
715 let bytes = upload.bytes();
716 let mut source_offset = 0_usize;
717 while source_offset < bytes.len() {
718 let remaining = bytes.len() - source_offset;
719 let chunk_bytes =
720 remaining.min(usize::try_from(maximum_chunk_bytes).map_err(|_| {
721 contract_failure(VNextError::InvalidExecutionPlan {
722 reason: "static staging budget exceeds host address space".to_owned(),
723 })
724 })?);
725 let chunk_bytes = chunk_bytes - chunk_bytes % element_bytes as usize;
726 if chunk_bytes == 0 {
727 return Err(contract_failure(VNextError::InvalidExecutionPlan {
728 reason: format!(
729 "component `{}` has a partial trailing element",
730 placement.component_id
731 ),
732 }));
733 }
734 let chunk_bytes_u64 = chunk_bytes as u64;
735 if !pending.is_empty()
736 && (pending.len() == policy.maximum_commands_per_batch()
737 || pending_staging_bytes
738 .checked_add(chunk_bytes_u64)
739 .is_none_or(|bytes| bytes > policy.maximum_staging_bytes()))
740 {
741 submission_wait_duration += submit_pending(
742 &runtime,
743 &mut stream,
744 &mut pending,
745 &mut pending_staging_bytes,
746 )?;
747 submission_batch_count += 1;
748 }
749 let source_end = source_offset + chunk_bytes;
750 let destination_offset = placement
751 .offset_bytes
752 .checked_add(source_offset as u64)
753 .ok_or_else(|| {
754 contract_failure(VNextError::InvalidExecutionPlan {
755 reason: "static upload destination offset overflows".to_owned(),
756 })
757 })?;
758 let layout =
759 HostTransferLayout::new(upload.element_type(), chunk_bytes_u64 / element_bytes)
760 .map_err(contract_failure)?;
761 let encode_started = Instant::now();
762 let command = with_static_buffer(transaction, &placement.resource_id, |buffer| {
763 runtime.encode_upload(
764 &bytes[source_offset..source_end],
765 layout,
766 buffer,
767 destination_offset,
768 )
769 })
770 .map_err(|error| {
771 runtime_or_contract_failure(&runtime, error, "static_upload_encode")
772 })?;
773 device_encode_duration += encode_started.elapsed();
774 pending.push(command);
775 pending_staging_bytes += chunk_bytes_u64;
776 upload_command_count += 1;
777 source_offset = source_end;
778 }
779 uploaded_component_count += 1;
780 uploaded_bytes = uploaded_bytes
781 .checked_add(placement.length_bytes)
782 .ok_or_else(|| {
783 contract_failure(VNextError::InvalidExecutionPlan {
784 reason: "static initialization uploaded bytes overflow u64".to_owned(),
785 })
786 })?;
787 }
788 }
789
790 if !pending.is_empty() {
791 submission_wait_duration += submit_pending(
792 &runtime,
793 &mut stream,
794 &mut pending,
795 &mut pending_staging_bytes,
796 )?;
797 submission_batch_count += 1;
798 }
799 if let Some(import) = weight_import {
800 let seal_started = Instant::now();
801 import.seal().map_err(|error| {
802 InitializationStepFailure::Quiescent(device_failure(
803 &runtime,
804 &error,
805 "static_weight_import_seal",
806 ))
807 })?;
808 import_seal_duration += seal_started.elapsed();
809 }
810 Ok(StaticInitializationReceipt {
811 initialized_resource_count: plan.payload().memory().static_allocations().len(),
812 uploaded_component_count,
813 uploaded_bytes,
814 imported_component_count,
815 imported_bytes,
816 transformed_component_count,
817 transformed_bytes,
818 transform_command_count,
819 upload_command_count,
820 submission_batch_count,
821 total_duration_us: duration_us(initialization_started.elapsed()),
822 setup_duration_us: duration_us(setup_duration),
823 source_materialization_duration_us: duration_us(source_materialization_duration),
824 device_encode_duration_us: duration_us(device_encode_duration),
825 device_import_duration_us: duration_us(device_import_duration),
826 device_transform_encode_duration_us: duration_us(device_transform_encode_duration),
827 submission_wait_duration_us: duration_us(submission_wait_duration),
828 import_seal_duration_us: duration_us(import_seal_duration),
829 slowest_component_id,
830 slowest_component_materialization_duration_us: duration_us(
831 slowest_component_materialization_duration,
832 ),
833 source_files,
834 })
835}
836
837fn submit_pending<R>(
838 runtime: &Arc<R>,
839 stream: &mut Option<R::Stream>,
840 pending: &mut Vec<R::Command>,
841 pending_staging_bytes: &mut u64,
842) -> Result<Duration, InitializationStepFailure<R>>
843where
844 R: DeviceRuntime,
845{
846 let started = Instant::now();
847 debug_assert!(!pending.is_empty());
848 let commands = std::mem::take(pending);
849 *pending_staging_bytes = 0;
850 let mut batch = DeviceCommandBatch::with_capacity(commands.len());
851 for command in commands {
852 batch.push_initialization(command);
853 }
854 let submitted = catch_unwind(AssertUnwindSafe(|| {
855 runtime.submit(
856 stream
857 .as_mut()
858 .expect("static initialization owns its stream"),
859 batch,
860 )
861 }));
862 let fence = match submitted {
863 Ok(Ok(fence)) => fence,
864 Ok(Err(not_submitted)) => {
865 return Err(InitializationStepFailure::Quiescent(device_failure(
866 runtime,
867 not_submitted.error(),
868 "static_submit_not_submitted",
869 )))
870 }
871 Err(payload) => {
872 return Err(InitializationStepFailure::Indeterminate {
873 failure: portable_failure(
874 FailureDomain::Device,
875 "static_submit_indeterminate",
876 panic_message(payload),
877 false,
878 ),
879 recovery: StaticInitializationRecovery {
880 stream: stream
881 .take()
882 .expect("static initialization owns its stream"),
883 fence: None,
884 },
885 })
886 }
887 };
888 let waited = catch_unwind(AssertUnwindSafe(|| runtime.wait_fence(&fence)));
889 match waited {
890 Ok(Ok(receipt)) => match receipt.into_parts().0 {
891 DeviceTerminal::Succeeded => Ok(started.elapsed()),
892 DeviceTerminal::FailedButQuiescent(error) => Err(InitializationStepFailure::Quiescent(
893 device_failure(runtime, &error, "static_fence_failed"),
894 )),
895 },
896 Ok(Err(indeterminate)) => Err(InitializationStepFailure::Indeterminate {
897 failure: device_failure(runtime, indeterminate.error(), "static_fence_indeterminate"),
898 recovery: StaticInitializationRecovery {
899 stream: stream
900 .take()
901 .expect("static initialization owns its stream"),
902 fence: Some(fence),
903 },
904 }),
905 Err(payload) => Err(InitializationStepFailure::Indeterminate {
906 failure: portable_failure(
907 FailureDomain::Device,
908 "static_fence_wait_panic",
909 panic_message(payload),
910 false,
911 ),
912 recovery: StaticInitializationRecovery {
913 stream: stream
914 .take()
915 .expect("static initialization owns its stream"),
916 fence: Some(fence),
917 },
918 }),
919 }
920}
921
922fn duration_us(duration: Duration) -> u64 {
923 u64::try_from(duration.as_micros()).unwrap_or(u64::MAX)
924}
925
926fn preflight_transaction<D>(
927 transaction: &ResourceTransaction<D, TransactionCommitted>,
928 family: &PreparedModelFamily,
929 plan: &ExecutionPlan,
930) -> Result<(), VNextError>
931where
932 D: ResourceTransactionDriver,
933{
934 let payload = plan.payload();
935 let admission = transaction.admission();
936 payload
937 .execution_weights()
938 .validate_against_family(family)?;
939 if payload.family_id() != family.family_id()
940 || payload.prepared_family_fingerprint() != family.fingerprint()?
941 || admission.plan_id() != payload.plan_id()
942 || admission.plan_hash() != plan.plan_hash()
943 || admission.device_id() != payload.device_id()
944 || admission.device_runtime_implementation_fingerprint()
945 != payload.device_runtime_implementation_fingerprint()
946 || transaction.lease().plan_static_entries().count()
947 != payload.memory().static_allocations().len()
948 {
949 return Err(VNextError::InvalidExecutionPlan {
950 reason: "static initialization family, plan, admission, runtime, or lease differs"
951 .to_owned(),
952 });
953 }
954 Ok(())
955}
956
957fn weight_placements(
958 family: &PreparedModelFamily,
959 plan: &ExecutionPlan,
960) -> Result<BTreeMap<WeightId, WeightPlacement>, VNextError> {
961 plan.payload()
962 .execution_weights()
963 .validate_against_family(family)?;
964 let execution_weight_schema = plan.payload().execution_weights().schema();
965 let schema = execution_weight_schema
966 .components
967 .iter()
968 .map(|component| (&component.id, component))
969 .collect::<BTreeMap<_, _>>();
970 let allocations = plan
971 .payload()
972 .memory()
973 .static_allocations()
974 .iter()
975 .map(|allocation| (allocation.resource_id(), allocation))
976 .collect::<BTreeMap<_, _>>();
977 let mut placements = BTreeMap::new();
978 for node in plan.payload().nodes() {
979 for binding in node
980 .values()
981 .iter()
982 .filter(|binding| binding.usage() == BufferUsage::Weights)
983 {
984 for resolved in binding.storage().components() {
985 let component_id =
986 resolved
987 .component_id()
988 .ok_or_else(|| VNextError::InvalidExecutionPlan {
989 reason: format!(
990 "weight resource `{}` lacks a physical component identity",
991 resolved.resource_id()
992 ),
993 })?;
994 let component =
995 schema
996 .get(component_id)
997 .ok_or_else(|| VNextError::InvalidExecutionPlan {
998 reason: format!("plan binds unknown weight component `{component_id}`"),
999 })?;
1000 let placement = WeightPlacement {
1001 component_id: component_id.clone(),
1002 resource_id: resolved.resource_id().clone(),
1003 offset_bytes: resolved.offset_bytes(),
1004 length_bytes: resolved.length_bytes(),
1005 element_type: resolved.element_type(),
1006 };
1007 if placement.length_bytes != component.physical_bytes()?
1008 || placement.element_type != component.physical_element_type()
1009 {
1010 return Err(VNextError::InvalidExecutionPlan {
1011 reason: format!(
1012 "weight component `{component_id}` placement differs from its physical schema"
1013 ),
1014 });
1015 }
1016 match placements.get(component_id) {
1017 Some(existing) if existing != &placement => {
1018 return Err(VNextError::InvalidExecutionPlan {
1019 reason: format!(
1020 "weight component `{component_id}` has inconsistent placements"
1021 ),
1022 })
1023 }
1024 Some(_) => {}
1025 None => {
1026 placements.insert(component_id.clone(), placement);
1027 }
1028 }
1029 }
1030 }
1031 }
1032 for component in &execution_weight_schema.components {
1033 if component.required && !placements.contains_key(&component.id) {
1034 return Err(VNextError::InvalidExecutionPlan {
1035 reason: format!(
1036 "required weight component `{}` has no plan placement",
1037 component.id
1038 ),
1039 });
1040 }
1041 }
1042 let mut ranges = BTreeMap::<ResourceId, Vec<(u64, u64, WeightId)>>::new();
1043 for placement in placements.values() {
1044 let allocation = allocations.get(&placement.resource_id).ok_or_else(|| {
1045 VNextError::InvalidExecutionPlan {
1046 reason: format!(
1047 "weight component `{}` references a non-static resource",
1048 placement.component_id
1049 ),
1050 }
1051 })?;
1052 let end = placement
1053 .offset_bytes
1054 .checked_add(placement.length_bytes)
1055 .ok_or_else(|| VNextError::InvalidExecutionPlan {
1056 reason: "weight placement range overflows u64".to_owned(),
1057 })?;
1058 if allocation.usage() != BufferUsage::Weights
1059 || allocation.element_type() != placement.element_type
1060 || end > allocation.size_bytes()
1061 {
1062 return Err(VNextError::InvalidExecutionPlan {
1063 reason: format!(
1064 "weight component `{}` placement exceeds or differs from its allocation",
1065 placement.component_id
1066 ),
1067 });
1068 }
1069 ranges
1070 .entry(placement.resource_id.clone())
1071 .or_default()
1072 .push((placement.offset_bytes, end, placement.component_id.clone()));
1073 }
1074 for (resource_id, ranges) in &mut ranges {
1075 ranges.sort();
1076 if ranges.windows(2).any(|pair| pair[0].1 > pair[1].0) {
1077 return Err(VNextError::InvalidExecutionPlan {
1078 reason: format!("weight placements overlap in resource `{resource_id}`"),
1079 });
1080 }
1081 }
1082 if allocations
1083 .values()
1084 .filter(|allocation| allocation.usage() == BufferUsage::Weights)
1085 .any(|allocation| !ranges.contains_key(allocation.resource_id()))
1086 {
1087 return Err(VNextError::InvalidExecutionPlan {
1088 reason: "a static weight allocation has no schema component placement".to_owned(),
1089 });
1090 }
1091 Ok(placements)
1092}
1093
1094fn prepare_uploads<'source>(
1095 family: &PreparedModelFamily,
1096 plan: &ExecutionPlan,
1097 source: &'source dyn WeightComponentSource,
1098 components: &[&WeightComponentSpec],
1099 placements: &BTreeMap<WeightId, WeightPlacement>,
1100) -> Result<Vec<WeightComponentPayload<'source>>, VNextError> {
1101 let payloads = plan.materialize_weight_components(family, source, components)?;
1102 for (component, payload) in components.iter().zip(&payloads) {
1103 let placement =
1104 placements
1105 .get(&component.id)
1106 .ok_or_else(|| VNextError::InvalidExecutionPlan {
1107 reason: format!(
1108 "execution weight component `{}` has no selected placement",
1109 component.id
1110 ),
1111 })?;
1112 if payload.component_id() != &placement.component_id
1113 || payload.element_type() != placement.element_type
1114 || payload.bytes().len() as u64 != placement.length_bytes
1115 {
1116 return Err(VNextError::InvalidExecutionPlan {
1117 reason: format!(
1118 "weight source payload for `{}` differs from its selected placement",
1119 placement.component_id
1120 ),
1121 });
1122 }
1123 }
1124 Ok(payloads)
1125}
1126
1127fn prepare_transform_sources<'source>(
1128 family: &PreparedModelFamily,
1129 source: &'source dyn WeightComponentSource,
1130 transform: &StaticWeightTransformPlan,
1131) -> Result<Vec<WeightComponentSegments<'source>>, VNextError> {
1132 transform
1133 .source_component_ids()
1134 .into_iter()
1135 .map(|source_id| {
1136 let component_index = family
1137 .weight_schema()
1138 .components
1139 .binary_search_by(|component| component.id.cmp(source_id))
1140 .map_err(|_| VNextError::InvalidExecutionPlan {
1141 reason: format!(
1142 "static weight transform references unknown source component `{source_id}`"
1143 ),
1144 })?;
1145 let component = &family.weight_schema().components[component_index];
1146 let segments = source.component_segments(component)?;
1147 if segments.component_id() != &component.id
1148 || segments.external_names() != component.external_names.as_slice()
1149 || segments.dimensions() != component.dimensions.as_slice()
1150 || segments.element_type() != component.physical_element_type()
1151 || segments.total_bytes() != component.physical_bytes()?
1152 {
1153 return Err(VNextError::InvalidExecutionPlan {
1154 reason: format!(
1155 "static weight transform source segments for `{}` differ from the trusted source schema",
1156 component.id
1157 ),
1158 });
1159 }
1160 Ok(segments)
1161 })
1162 .collect()
1163}
1164
1165#[allow(clippy::too_many_arguments)]
1166fn encode_required_weight_transform<'source, D>(
1167 transaction: &ResourceTransaction<D, TransactionCommitted>,
1168 runtime: &Arc<D::Runtime>,
1169 transform: &StaticWeightTransformPlan,
1170 sources: &[WeightComponentSegments<'source>],
1171 components: &[&WeightComponentSpec],
1172 placements: &BTreeMap<WeightId, WeightPlacement>,
1173 scratch_resource_id: &ResourceId,
1174) -> Result<
1175 <D::Runtime as DeviceRuntime>::Command,
1176 StaticBufferAccessError<<D::Runtime as DeviceRuntime>::Error>,
1177>
1178where
1179 D: ResourceTransactionDriver,
1180{
1181 let [packed_values_id, scales_id] = transform.execution_component_ids();
1182 let packed_component = components
1183 .iter()
1184 .copied()
1185 .find(|component| &component.id == packed_values_id)
1186 .ok_or_else(|| {
1187 StaticBufferAccessError::Contract(VNextError::InvalidExecutionPlan {
1188 reason: "static weight transform packed output is absent from its component group"
1189 .to_owned(),
1190 })
1191 })?;
1192 let scales_component = components
1193 .iter()
1194 .copied()
1195 .find(|component| &component.id == scales_id)
1196 .ok_or_else(|| {
1197 StaticBufferAccessError::Contract(VNextError::InvalidExecutionPlan {
1198 reason: "static weight transform scale output is absent from its component group"
1199 .to_owned(),
1200 })
1201 })?;
1202 let packed_placement = placements.get(packed_values_id).ok_or_else(|| {
1203 StaticBufferAccessError::Contract(VNextError::InvalidExecutionPlan {
1204 reason: "static weight transform packed output has no placement".to_owned(),
1205 })
1206 })?;
1207 let scales_placement = placements.get(scales_id).ok_or_else(|| {
1208 StaticBufferAccessError::Contract(VNextError::InvalidExecutionPlan {
1209 reason: "static weight transform scale output has no placement".to_owned(),
1210 })
1211 })?;
1212 let lease = transaction.lease();
1213 let packed_entry = lease
1214 .plan_static_entries()
1215 .find(|entry| entry.resource_id() == &packed_placement.resource_id)
1216 .ok_or_else(|| {
1217 StaticBufferAccessError::Contract(VNextError::InvalidExecutionPlan {
1218 reason: format!(
1219 "static lease lacks transform destination `{}`",
1220 packed_placement.resource_id
1221 ),
1222 })
1223 })?;
1224 let scales_entry = lease
1225 .plan_static_entries()
1226 .find(|entry| entry.resource_id() == &scales_placement.resource_id)
1227 .ok_or_else(|| {
1228 StaticBufferAccessError::Contract(VNextError::InvalidExecutionPlan {
1229 reason: format!(
1230 "static lease lacks transform destination `{}`",
1231 scales_placement.resource_id
1232 ),
1233 })
1234 })?;
1235 let scratch_entry = lease
1236 .plan_static_entries()
1237 .find(|entry| entry.resource_id() == scratch_resource_id)
1238 .ok_or_else(|| {
1239 StaticBufferAccessError::Contract(VNextError::InvalidExecutionPlan {
1240 reason: format!("static lease lacks transform scratch `{scratch_resource_id}`"),
1241 })
1242 })?;
1243 let packed_view = lease
1244 .view(&packed_placement.resource_id, packed_entry.generation())
1245 .map_err(StaticBufferAccessError::Contract)?;
1246 let scales_view = lease
1247 .view(&scales_placement.resource_id, scales_entry.generation())
1248 .map_err(StaticBufferAccessError::Contract)?;
1249 let scratch_view = lease
1250 .view(scratch_resource_id, scratch_entry.generation())
1251 .map_err(StaticBufferAccessError::Contract)?;
1252 let destinations = [
1253 StaticWeightTransformDestination::new(
1254 packed_component,
1255 packed_view.buffer(),
1256 packed_placement.offset_bytes,
1257 ),
1258 StaticWeightTransformDestination::new(
1259 scales_component,
1260 scales_view.buffer(),
1261 scales_placement.offset_bytes,
1262 ),
1263 ];
1264 let request =
1265 StaticWeightTransformRequest::new(transform, sources, &destinations, scratch_view.buffer());
1266 match runtime.encode_static_weight_transform(request) {
1267 Some(Ok(command)) => Ok(command),
1268 Some(Err(error)) => Err(StaticBufferAccessError::Runtime(error)),
1269 None => Err(StaticBufferAccessError::Contract(
1270 VNextError::InvalidExecutionPlan {
1271 reason: "device runtime does not support the required static weight transform"
1272 .to_owned(),
1273 },
1274 )),
1275 }
1276}
1277
1278enum StaticBufferAccessError<E> {
1279 Contract(VNextError),
1280 Runtime(E),
1281}
1282
1283fn with_static_buffer<D, T>(
1284 transaction: &ResourceTransaction<D, TransactionCommitted>,
1285 resource_id: &ResourceId,
1286 action: impl FnOnce(&D::Buffer) -> Result<T, <D::Runtime as DeviceRuntime>::Error>,
1287) -> Result<T, StaticBufferAccessError<<D::Runtime as DeviceRuntime>::Error>>
1288where
1289 D: ResourceTransactionDriver,
1290{
1291 let lease = transaction.lease();
1292 let entry = lease
1293 .plan_static_entries()
1294 .find(|entry| entry.resource_id() == resource_id)
1295 .ok_or_else(|| {
1296 StaticBufferAccessError::Contract(VNextError::InvalidExecutionPlan {
1297 reason: format!("static lease lacks resource `{resource_id}`"),
1298 })
1299 })?;
1300 let view = lease
1301 .view(resource_id, entry.generation())
1302 .map_err(StaticBufferAccessError::Contract)?;
1303 action(view.buffer()).map_err(StaticBufferAccessError::Runtime)
1304}
1305
1306fn runtime_or_contract_failure<R>(
1307 runtime: &Arc<R>,
1308 error: StaticBufferAccessError<R::Error>,
1309 code: &'static str,
1310) -> InitializationStepFailure<R>
1311where
1312 R: DeviceRuntime,
1313{
1314 InitializationStepFailure::Quiescent(match error {
1315 StaticBufferAccessError::Contract(error) => resource_failure(code, error),
1316 StaticBufferAccessError::Runtime(error) => device_failure(runtime, &error, code),
1317 })
1318}
1319
1320fn contract_failure<R>(error: VNextError) -> InitializationStepFailure<R>
1321where
1322 R: DeviceRuntime,
1323{
1324 InitializationStepFailure::Quiescent(resource_failure("static_contract", error))
1325}
1326
1327fn resource_failure(code: &'static str, error: impl fmt::Display) -> FailureEnvelope {
1328 portable_failure(FailureDomain::Resource, code, error, false)
1329}
1330
1331fn device_failure<R>(
1332 runtime: &Arc<R>,
1333 error: &R::Error,
1334 fallback_code: &'static str,
1335) -> FailureEnvelope
1336where
1337 R: DeviceRuntime,
1338{
1339 match catch_unwind(AssertUnwindSafe(|| runtime.describe_error(error))) {
1340 Ok(Ok(report)) => portable_failure(
1341 FailureDomain::Device,
1342 report.code(),
1343 report.message(),
1344 report.retryable(),
1345 ),
1346 Ok(Err(classification)) => portable_failure(
1347 FailureDomain::Device,
1348 fallback_code,
1349 format!("{error}; error classification failed: {classification}"),
1350 false,
1351 ),
1352 Err(payload) => portable_failure(
1353 FailureDomain::Device,
1354 fallback_code,
1355 format!(
1356 "{error}; error classification panicked: {}",
1357 panic_message(payload)
1358 ),
1359 false,
1360 ),
1361 }
1362}
1363
1364fn portable_failure(
1365 domain: FailureDomain,
1366 code: impl Into<String>,
1367 message: impl fmt::Display,
1368 retryable: bool,
1369) -> FailureEnvelope {
1370 let mut code = code.into();
1371 code.retain(|character| {
1372 character.is_ascii_alphanumeric() || matches!(character, '.' | '_' | '-')
1373 });
1374 code.truncate(64);
1375 if code.is_empty() {
1376 code.push_str("static_initialization");
1377 }
1378 let mut message = message
1379 .to_string()
1380 .chars()
1381 .filter(|character| !character.is_control() || matches!(character, '\n' | '\t'))
1382 .take(1024)
1383 .collect::<String>();
1384 if message.trim().is_empty() {
1385 message.push_str("static initialization failed");
1386 }
1387 FailureEnvelope::new(domain, code, message, retryable)
1388 .expect("static initialization failure metadata is bounded and portable")
1389}
1390
1391fn panic_message(payload: Box<dyn Any + Send>) -> String {
1392 if let Some(message) = payload.downcast_ref::<&str>() {
1393 (*message).to_owned()
1394 } else if let Some(message) = payload.downcast_ref::<String>() {
1395 message.clone()
1396 } else {
1397 "device runtime panicked during static initialization submission".to_owned()
1398 }
1399}