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