1use crate::numeric::WGPU_NUMERIC;
4use smallvec::SmallVec;
5use std::cell::RefCell;
6use std::sync::Arc;
7use std::time::Instant;
8use vyre_driver::{
9 BackendError, CompiledPipeline, DispatchConfig, OutputBuffers, Resource, VyreBackend,
10};
11use vyre_foundation::ir::Program;
12
13use vyre_runtime::megakernel::io::{
14 try_encode_empty_io_queue_into, validate_io_queue_bytes, IO_SLOT_COUNT,
15};
16use vyre_runtime::megakernel::{
17 build_program_sharded_once_slots_control_report_shared, build_scallop_lineage_with_scratch,
18 plan_compact_fusion_into, try_prune_redundant_work_items_with_scratch_into,
19 CompactFusionPlanningScratch, CrossArmRedundancy, Megakernel, MegakernelConfig,
20 MegakernelDispatch, MegakernelLaunchRecommendation, MegakernelReport, MegakernelTelemetry,
21 MegakernelWorkItem, RedundantWorkItemPruneScratch, IO_SLOT_WORDS,
22};
23
24#[cfg(feature = "megakernel-batch")]
25#[path = "megakernel/batch.rs"]
26pub mod batch;
27#[cfg(feature = "megakernel-batch")]
28#[path = "megakernel/dispatch_plan.rs"]
29pub mod dispatch_plan;
30#[cfg(feature = "megakernel-batch")]
31#[path = "megakernel/dispatcher.rs"]
32pub mod dispatcher;
33#[cfg(feature = "megakernel-batch")]
34#[path = "megakernel/pipeline_cache.rs"]
35pub(crate) mod pipeline_cache;
36#[cfg(feature = "megakernel-batch")]
37#[path = "megakernel/segmentation.rs"]
38pub mod segmentation;
39
40#[cfg(feature = "megakernel-batch")]
41pub use batch::{
42 queue_state_word, BatchFile, CombinedBatch, FileBatch, FileBatchRefreshReport, FileMetadata,
43 HitRecord, WorkTriple, FILE_METADATA_WORDS, HIT_RECORD_WORDS, QUEUE_STATE_WORDS,
44 WORK_TRIPLE_WORDS,
45};
46#[cfg(feature = "megakernel-batch")]
47pub use dispatch_plan::BatchDispatchPlan;
48#[cfg(feature = "megakernel-batch")]
49pub use dispatcher::{
50 wgpu_scan_batch_segmentation_evidence, BatchDispatchConfig, BatchDispatchReport,
51 BatchDispatchSummary, BatchDispatchTelemetry, BatchDispatcher, BatchHitWriter,
52 CombinedDispatchSummary, CombinedDispatcher, SegLenCalibration, SegLenMeasurement,
53 TransitionWidth, WgpuScanBatchSegmentationError, WgpuScanBatchSegmentationEvidence,
54 WgpuScanBatchSegmentationRequest, DEFAULT_SEG_LEN_CANDIDATES,
55 WGPU_SCAN_BATCH_SEGMENTATION_SCHEMA_VERSION,
56};
57
58thread_local! {
59 static DISPATCH_SCRATCH: RefCell<DispatchScratch> = RefCell::new(DispatchScratch::default());
60}
61
62const MAX_INLINE_LINEAGE_ITEMS: usize = 256;
63
64#[derive(Default)]
65struct DispatchScratch {
66 io_queue_bytes: Vec<u8>,
67 control_bytes: Vec<u8>,
68 ring_words: Vec<u32>,
69 debug_log_bytes: Vec<u8>,
70 fusion: CompactFusionPlanningScratch,
71 lineage_state: Vec<u32>,
72 lineage_next: Vec<u32>,
73 lineage_changed: [u32; 1],
74 deduped_items: Vec<MegakernelWorkItem>,
75 dedupe: RedundantWorkItemPruneScratch,
76 compiled: Option<CompiledMegakernelPipeline>,
77 resident: Option<ResidentMegakernelBuffers>,
78 outputs: OutputBuffers,
79}
80
81struct CompiledMegakernelPipeline {
82 backend_id: &'static str,
83 backend_version: &'static str,
84 workgroup_size_x: u32,
85 slot_count: u32,
86 dispatch_config: DispatchConfig,
87 program: Arc<Program>,
88 pipeline: Arc<dyn CompiledPipeline>,
89}
90
91struct ResidentMegakernelBuffers {
92 backend_id: &'static str,
93 backend_version: &'static str,
94 workgroup_size_x: u32,
95 slot_count: u32,
96 input_lens: [usize; 4],
97 resources: Vec<Resource>,
98}
99
100enum IoQueueInput<'a> {
101 Scratch,
102 Borrowed(&'a [u8]),
103}
104
105pub struct WgpuMegakernelDispatcher<'a> {
107 backend: &'a dyn VyreBackend,
108}
109
110impl<'a> WgpuMegakernelDispatcher<'a> {
111 #[must_use]
113 pub fn new(backend: &'a dyn VyreBackend) -> Self {
114 Self { backend }
115 }
116
117 pub fn dispatch_megakernel_bytes(
124 &self,
125 work_queue_bytes: &[u8],
126 config: &MegakernelConfig,
127 ) -> Result<MegakernelReport, BackendError> {
128 if work_queue_bytes.len() % std::mem::size_of::<MegakernelWorkItem>() != 0 {
129 return Err(BackendError::new(format!(
130 "megakernel work queue has {} bytes, which is not a multiple of sizeof(MegakernelWorkItem)={}. Fix: encode whole MegakernelWorkItem records before dispatch.",
131 work_queue_bytes.len(),
132 std::mem::size_of::<MegakernelWorkItem>()
133 )));
134 }
135 let work_items = bytemuck::try_cast_slice::<u8, MegakernelWorkItem>(work_queue_bytes).map_err(|err| {
136 BackendError::new(format!(
137 "megakernel work queue bytes are not aligned as MegakernelWorkItem records: {err}. Fix: allocate or copy the queue into aligned MegakernelWorkItem storage before dispatch."
138 ))
139 })?;
140 self.dispatch_megakernel(work_items, config)
141 }
142
143 pub fn dispatch_megakernel(
145 &self,
146 work_items: &[MegakernelWorkItem],
147 config: &MegakernelConfig,
148 ) -> Result<MegakernelReport, BackendError> {
149 config.validate()?;
150
151 if work_items.is_empty() {
152 return Ok(MegakernelReport::default());
153 }
154
155 DISPATCH_SCRATCH.with(|scratch| {
156 let mut scratch = scratch.borrow_mut();
157 ensure_empty_io_queue_bytes(&mut scratch.io_queue_bytes)?;
158 self.dispatch_megakernel_with_io_queue_ref(
159 work_items,
160 config,
161 IoQueueInput::Scratch,
162 &mut scratch,
163 )
164 })
165 }
166
167 pub fn dispatch_megakernel_with_io_queue(
173 &self,
174 work_items: &[MegakernelWorkItem],
175 config: &MegakernelConfig,
176 io_queue_bytes: Vec<u8>,
177 ) -> Result<MegakernelReport, BackendError> {
178 DISPATCH_SCRATCH.with(|scratch| {
179 let mut scratch = scratch.borrow_mut();
180 self.dispatch_megakernel_with_io_queue_ref(
181 work_items,
182 config,
183 IoQueueInput::Borrowed(io_queue_bytes.as_slice()),
184 &mut scratch,
185 )
186 })
187 }
188
189 fn dispatch_megakernel_with_io_queue_ref(
190 &self,
191 work_items: &[MegakernelWorkItem],
192 config: &MegakernelConfig,
193 io_queue: IoQueueInput<'_>,
194 scratch: &mut DispatchScratch,
195 ) -> Result<MegakernelReport, BackendError> {
196 config.validate()?;
197 let io_queue_bytes = match io_queue {
198 IoQueueInput::Scratch => scratch.io_queue_bytes.as_slice(),
199 IoQueueInput::Borrowed(bytes) => bytes,
200 };
201 validate_io_queue_bytes(io_queue_bytes).map_err(|e| BackendError::new(e.to_string()))?;
202
203 let initial_item_count = work_items.len();
204 if initial_item_count == 0 {
205 return Ok(MegakernelReport::default());
206 }
207
208 let plan_start = Instant::now();
209 let redundancy = try_prune_redundant_work_items_with_scratch_into(
210 work_items,
211 &mut scratch.deduped_items,
212 &mut scratch.dedupe,
213 )
214 .map_err(|error| BackendError::new(error.to_string()))?;
215 let planning_items = if redundancy.is_empty() {
216 work_items
217 } else {
218 scratch.deduped_items.as_slice()
219 };
220
221 let track_lineage = should_track_lineage(planning_items.len());
222 if track_lineage {
223 let _fusion_plan = plan_compact_fusion_into(planning_items, &mut scratch.fusion);
224 } else {
225 let empty_fusion_plan = plan_compact_fusion_into(&[], &mut scratch.fusion);
226 debug_assert!(
227 empty_fusion_plan.is_empty(),
228 "empty megakernel fusion planning input must produce an empty plan"
229 );
230 }
231 let dispatch_items = planning_items;
232 let item_count = dispatch_items.len();
233 let queue_plan_ns = nanos_u64(plan_start.elapsed().as_nanos())?;
234
235 let queue_len = u32::try_from(item_count).map_err(|_| {
236 BackendError::new(
237 "megakernel work queue length exceeds u32::MAX. Fix: shard the queue before dispatch.",
238 )
239 })?;
240 let max_workgroup_size_x = self.backend.max_workgroup_size()[0];
241 if max_workgroup_size_x == 0 {
242 return Err(BackendError::new(format!(
243 "backend `{}` reported max_workgroup_size.x=0. Fix: use a backend that exposes real adapter limits before megakernel dispatch.",
244 self.backend.id()
245 )));
246 }
247 let launch = config.launch_recommendation(
248 queue_len,
249 max_workgroup_size_x,
250 self.backend.max_compute_workgroups_per_dimension(),
251 self.backend.max_compute_invocations_per_workgroup(),
252 )?;
253 let geometry = launch.geometry;
254
255 let publish_start = Instant::now();
256 let dispatch_config = geometry.dispatch_config(Some(config.max_wall_time));
257 let compiled_cache_hit = compiled_pipeline_cache_matches(
258 self.backend,
259 geometry.workgroup_size_x,
260 geometry.slot_count,
261 &dispatch_config,
262 &scratch.compiled,
263 );
264 let program = if compiled_cache_hit {
265 None
266 } else {
267 Some(build_program_sharded_once_slots_control_report_shared(
268 geometry.workgroup_size_x,
269 geometry.slot_count,
270 &[],
271 ))
272 };
273 let compiled = if compiled_cache_hit {
274 scratch
275 .compiled
276 .as_ref()
277 .map(|cached| cached.pipeline.as_ref())
278 } else {
279 let program = program.as_ref().ok_or_else(|| {
280 BackendError::new(
281 "megakernel cache miss had no Program to compile. Fix: build the sharded megakernel Program before compiling a new geometry."
282 .to_string(),
283 )
284 })?;
285 compiled_pipeline_for_geometry(
286 self.backend,
287 program.clone(),
288 geometry.workgroup_size_x,
289 geometry.slot_count,
290 &dispatch_config,
291 &mut scratch.compiled,
292 )?
293 };
294 Megakernel::encode_work_items_ring_words_into(
295 geometry.slot_count,
296 0,
297 dispatch_items,
298 &mut scratch.ring_words,
299 )
300 .map_err(|e| BackendError::new(e.to_string()))?;
301 ensure_control_bytes(&mut scratch.control_bytes)?;
302 ensure_empty_debug_log_bytes(&mut scratch.debug_log_bytes)?;
303 let queue_publish_ns = nanos_u64(publish_start.elapsed().as_nanos())?;
304
305 let start = Instant::now();
306 let inputs = [
307 scratch.control_bytes.as_slice(),
308 bytemuck::cast_slice(scratch.ring_words.as_slice()),
309 scratch.debug_log_bytes.as_slice(),
310 io_queue_bytes,
311 ];
312 let estimated_peak_device_bytes = megakernel_dispatch_peak_device_bytes(&inputs, launch)?;
313 enforce_megakernel_device_memory_budget(
314 estimated_peak_device_bytes,
315 launch.device_memory_budget_bytes,
316 )?;
317 let mut resident_allocations = 0;
318 let mut resident_input_cache_hit = false;
319 if let Some(compiled) = compiled {
320 resident_allocations = resident_megakernel_allocation_events(
321 self.backend,
322 geometry.workgroup_size_x,
323 geometry.slot_count,
324 &inputs,
325 &scratch.resident,
326 );
327 let input_lens = [
328 inputs[0].len(),
329 inputs[1].len(),
330 inputs[2].len(),
331 inputs[3].len(),
332 ];
333 resident_input_cache_hit = resident_megakernel_cache_matches(
334 self.backend,
335 geometry.workgroup_size_x,
336 geometry.slot_count,
337 input_lens,
338 &scratch.resident,
339 );
340 if let Some(resources) = ensure_resident_megakernel_buffers(
341 self.backend,
342 geometry.workgroup_size_x,
343 geometry.slot_count,
344 &inputs,
345 &mut scratch.resident,
346 )? {
347 compiled.dispatch_persistent_handles_into(
348 resources,
349 &dispatch_config,
350 &mut scratch.outputs,
351 )?;
352 } else {
353 compiled.dispatch_borrowed_into(&inputs, &dispatch_config, &mut scratch.outputs)?;
354 }
355 } else {
356 let program = program.ok_or_else(|| {
357 BackendError::new(
358 "megakernel cache-miss dispatch had no compiled pipeline and no Program. Fix: build the megakernel Program on every non-native-cache path.".to_string(),
359 )
360 })?;
361 self.backend.dispatch_borrowed_into(
362 program.as_ref(),
363 &inputs,
364 &dispatch_config,
365 &mut scratch.outputs,
366 )?;
367 }
368 let wall_time = start.elapsed();
369
370 let control_done_count = scratch.outputs.first().ok_or_else(|| {
371 BackendError::new(
372 "megakernel dispatch returned no control output buffer. Fix: backend must return the control buffer as output 0.",
373 )
374 })?;
375 let control_done_count = u64::from(
376 Megakernel::try_read_done_count(control_done_count)
377 .map_err(|error| BackendError::new(error.to_string()))?,
378 );
379 let slot_done_count = strict_done_ring_slots_from_outputs(&scratch.outputs, item_count)?;
380 let done_count = control_done_count.max(slot_done_count);
381
382 let lineage_start = Instant::now();
393 let region_lineage = if track_lineage {
394 build_scallop_lineage_with_scratch(
395 self.backend,
396 planning_items,
397 scratch.fusion.exchange_adj(),
398 planning_items.len(),
399 &mut scratch.lineage_state,
400 &mut scratch.lineage_next,
401 &mut scratch.lineage_changed,
402 config.max_wall_time,
403 )?
404 } else {
405 Vec::new()
406 };
407 let lineage_ns = nanos_u64(lineage_start.elapsed().as_nanos())?;
408
409 let redundant_items = retained_redundant_done_count(
410 work_items,
411 dispatch_items,
412 done_count,
413 item_count,
414 &redundancy,
415 );
416 let initial_item_count_u64 = u64::try_from(initial_item_count).map_err(|source| {
417 BackendError::new(format!(
418 "megakernel initial item count cannot fit u64: {source}. Fix: shard work items before dispatch."
419 ))
420 })?;
421 let logical_done_count = done_count.checked_add(redundant_items).ok_or_else(|| {
422 BackendError::new(
423 "megakernel logical done count overflowed u64. Fix: shard work items before dispatch.",
424 )
425 })?;
426 let bounded_done_count = logical_done_count.min(initial_item_count_u64);
427 let telemetry = megakernel_report_telemetry(
428 &inputs,
429 &scratch.outputs,
430 resident_allocations,
431 item_count,
432 geometry.slot_count,
433 geometry.covering_worker_groups(),
434 geometry.workgroup_size_x,
435 launch,
436 estimated_peak_device_bytes,
437 compiled_cache_hit,
438 resident_input_cache_hit,
439 )?;
440 Ok(MegakernelReport {
441 items_processed: bounded_done_count,
442 items_remaining: initial_item_count_u64 - bounded_done_count,
443 wall_time,
444 queue_plan_ns,
445 queue_publish_ns,
446 backend_dispatch_ns: nanos_u64(wall_time.as_nanos())?,
447 lineage_ns,
448 deduped_items: WGPU_NUMERIC.usize_to_u64(
449 redundancy.total_redundant_ops,
450 "megakernel redundant operation count",
451 )?,
452 published_items: WGPU_NUMERIC
453 .usize_to_u64(item_count, "megakernel published item count")?,
454 lineage_items: if track_lineage {
455 WGPU_NUMERIC.usize_to_u64(item_count, "megakernel lineage item count")?
456 } else {
457 0
458 },
459 telemetry,
460 region_lineage,
461 })
462 }
463}
464
465fn strict_done_ring_slots_from_outputs(
466 outputs: &[Vec<u8>],
467 item_count: usize,
468) -> Result<u64, BackendError> {
469 if item_count == 0 {
470 return Ok(0);
471 }
472 let item_count_u32 = u32::try_from(item_count).map_err(|source| {
473 BackendError::new(format!(
474 "megakernel item_count {item_count} cannot fit u32 for ring decode: {source}. Fix: shard megakernel dispatches before protocol ring sizing."
475 ))
476 })?;
477 let ring_bytes = Megakernel::ring_byte_len(item_count_u32).ok_or_else(|| {
478 BackendError::new(
479 "megakernel item_count ring byte length overflowed. Fix: shard megakernel dispatches before protocol ring sizing.".to_string(),
480 )
481 })?;
482 let mut saw_ring_output = false;
483 let mut max_done = 0_u64;
484 for bytes in outputs {
485 if bytes.len() < ring_bytes {
486 continue;
487 }
488 saw_ring_output = true;
489 let done = Megakernel::try_count_done_ring_slots(bytes, item_count)
490 .map_err(|source| BackendError::new(source.to_string()))?;
491 max_done = max_done.max(done);
492 }
493 if !saw_ring_output {
494 return Err(BackendError::new(format!(
495 "megakernel dispatch returned no output buffer large enough for {item_count} ring slot(s). Fix: backend output 0 must include control and at least one ring-sized status readback."
496 )));
497 }
498 Ok(max_done)
499}
500
501fn ensure_resident_megakernel_buffers<'a>(
502 backend: &dyn VyreBackend,
503 workgroup_size_x: u32,
504 slot_count: u32,
505 inputs: &[&[u8]; 4],
506 cache: &'a mut Option<ResidentMegakernelBuffers>,
507) -> Result<Option<&'a [Resource]>, BackendError> {
508 if backend.id() != "cuda" {
509 return Ok(None);
510 }
511
512 let input_lens = [
513 inputs[0].len(),
514 inputs[1].len(),
515 inputs[2].len(),
516 inputs[3].len(),
517 ];
518 let matches_cache =
519 resident_megakernel_cache_matches(backend, workgroup_size_x, slot_count, input_lens, cache);
520
521 if !matches_cache {
522 if let Some(old) = cache.take() {
523 for resource in old.resources {
524 backend.free_resident(resource)?;
525 }
526 }
527 let mut resources = Vec::with_capacity(inputs.len());
528 for input in inputs {
529 match backend.allocate_resident(input.len()) {
530 Ok(resource) => resources.push(resource),
531 Err(BackendError::UnsupportedFeature { name, backend }) => {
532 return Err(BackendError::UnsupportedFeature {
533 name: format!(
534 "CUDA resident megakernel input allocation required `{name}`"
535 ),
536 backend,
537 });
538 }
539 Err(error) => return Err(error),
540 }
541 }
542 if let Err(error) = refresh_resident_megakernel_inputs(backend, &resources, &inputs) {
543 return match release_resident_megakernel_resources(backend, resources) {
544 Ok(()) => Err(error),
545 Err(cleanup) => Err(BackendError::new(format!(
546 "megakernel resident input upload failed, and cleanup of newly allocated resident slots also failed. Upload error: {error}. Cleanup error: {cleanup}. Fix: inspect CUDA resident buffer ownership and stream state before retrying."
547 ))),
548 };
549 }
550 *cache = Some(ResidentMegakernelBuffers {
551 backend_id: backend.id(),
552 backend_version: backend.version(),
553 workgroup_size_x,
554 slot_count,
555 input_lens,
556 resources,
557 });
558 } else if let Some(resident) = cache.as_ref() {
559 refresh_resident_megakernel_inputs(backend, &resident.resources, &inputs)?;
560 }
561
562 Ok(cache.as_ref().map(|resident| resident.resources.as_slice()))
563}
564
565fn resident_megakernel_allocation_events(
566 backend: &dyn VyreBackend,
567 workgroup_size_x: u32,
568 slot_count: u32,
569 inputs: &[&[u8]; 4],
570 cache: &Option<ResidentMegakernelBuffers>,
571) -> u32 {
572 if backend.id() != "cuda" {
573 return 0;
574 }
575 let input_lens = [
576 inputs[0].len(),
577 inputs[1].len(),
578 inputs[2].len(),
579 inputs[3].len(),
580 ];
581 if resident_megakernel_cache_matches(backend, workgroup_size_x, slot_count, input_lens, cache) {
582 0
583 } else {
584 4
585 }
586}
587
588fn resident_megakernel_cache_matches(
589 backend: &dyn VyreBackend,
590 workgroup_size_x: u32,
591 slot_count: u32,
592 input_lens: [usize; 4],
593 cache: &Option<ResidentMegakernelBuffers>,
594) -> bool {
595 cache.as_ref().is_some_and(|resident| {
596 resident.backend_id == backend.id()
597 && resident.backend_version == backend.version()
598 && resident.workgroup_size_x == workgroup_size_x
599 && resident.slot_count == slot_count
600 && resident.input_lens == input_lens
601 })
602}
603
604fn refresh_resident_megakernel_inputs(
605 backend: &dyn VyreBackend,
606 resources: &[Resource],
607 inputs: &[&[u8]; 4],
608) -> Result<(), BackendError> {
609 let uploads = resident_input_upload_plan(resources, inputs)?;
610 backend.upload_resident_many(uploads.as_slice())
611}
612
613fn resident_input_upload_plan<'a>(
614 resources: &'a [Resource],
615 inputs: &'a [&[u8]; 4],
616) -> Result<SmallVec<[(&'a Resource, &'a [u8]); 4]>, BackendError> {
617 if resources.len() != inputs.len() {
618 return Err(BackendError::new(format!(
619 "megakernel resident input refresh expected {} resident slot(s) for {} input buffer(s). Fix: rebuild resident megakernel resources when the ABI input count changes.",
620 inputs.len(),
621 resources.len()
622 )));
623 }
624 let mut uploads = SmallVec::with_capacity(inputs.len());
625 for (resource, input) in resources.iter().zip(inputs.iter()) {
626 uploads.push((resource, *input));
627 }
628 Ok(uploads)
629}
630
631fn release_resident_megakernel_resources(
632 backend: &dyn VyreBackend,
633 resources: Vec<Resource>,
634) -> Result<(), BackendError> {
635 let mut first_error = None;
636 for resource in resources {
637 if let Err(error) = backend.free_resident(resource) {
638 if first_error.is_none() {
639 first_error = Some(error);
640 }
641 }
642 }
643 if let Some(error) = first_error {
644 Err(error)
645 } else {
646 Ok(())
647 }
648}
649
650fn megakernel_report_telemetry(
651 inputs: &[&[u8]; 4],
652 outputs: &OutputBuffers,
653 resident_allocations: u32,
654 item_count: usize,
655 slot_count: u32,
656 worker_groups: u32,
657 workgroup_size_x: u32,
658 launch: MegakernelLaunchRecommendation,
659 estimated_peak_device_bytes: u64,
660 compiled_pipeline_cache_hit: bool,
661 resident_input_cache_hit: bool,
662) -> Result<MegakernelTelemetry, BackendError> {
663 let mut bytes_uploaded = 0u64;
664 for input in inputs {
665 let input_len = u64::try_from(input.len()).map_err(|source| {
666 BackendError::new(format!(
667 "megakernel telemetry input length cannot fit u64: {source}. Fix: shard input buffers before dispatch."
668 ))
669 })?;
670 bytes_uploaded = bytes_uploaded.checked_add(input_len).ok_or_else(|| {
671 BackendError::new(
672 "megakernel telemetry uploaded-byte total overflowed u64. Fix: shard input buffers before dispatch.",
673 )
674 })?;
675 }
676 let mut bytes_read_back = 0u64;
677 for output in outputs {
678 let output_len = u64::try_from(output.len()).map_err(|source| {
679 BackendError::new(format!(
680 "megakernel telemetry output length cannot fit u64: {source}. Fix: shard output buffers before dispatch."
681 ))
682 })?;
683 bytes_read_back = bytes_read_back.checked_add(output_len).ok_or_else(|| {
684 BackendError::new(
685 "megakernel telemetry readback-byte total overflowed u64. Fix: shard output buffers before dispatch.",
686 )
687 })?;
688 }
689 let bytes_moved = bytes_uploaded.checked_add(bytes_read_back).ok_or_else(|| {
690 BackendError::new(
691 "megakernel telemetry moved-byte total overflowed u64. Fix: shard dispatch buffers before dispatch.",
692 )
693 })?;
694 Ok(MegakernelTelemetry {
695 bytes_uploaded,
696 bytes_read_back,
697 bytes_moved,
698 resident_allocations,
699 kernel_launches: 1,
700 sync_points: 1,
701 occupancy_proxy_bps: occupancy_proxy_bps(item_count, worker_groups, workgroup_size_x)?,
702 frontier_density_bps: density_bps(
703 WGPU_NUMERIC.usize_to_u64(item_count, "megakernel frontier item count")?,
704 u64::from(slot_count.max(1)),
705 ),
706 readback_buffers: u32::try_from(outputs.len()).map_err(|source| {
707 BackendError::new(format!(
708 "megakernel readback buffer count cannot fit u32: {source}. Fix: shard output buffers before telemetry reporting."
709 ))
710 })?,
711 compiled_pipeline_cache_hit,
712 resident_input_cache_hit,
713 topology: launch.topology,
714 pressure: launch.pressure,
715 execution_mode: launch.execution_mode,
716 hit_capacity: launch.hit_capacity,
717 estimated_peak_device_bytes,
718 device_memory_budget_bytes: launch.device_memory_budget_bytes,
719 })
720}
721
722fn megakernel_dispatch_peak_device_bytes(
723 inputs: &[&[u8]; 4],
724 launch: MegakernelLaunchRecommendation,
725) -> Result<u64, BackendError> {
726 let mut abi_input_bytes = 0u64;
727 for input in inputs {
728 let input_len = u64::try_from(input.len()).map_err(|source| {
729 BackendError::new(format!(
730 "megakernel ABI input length cannot fit u64: {source}. Fix: shard ABI input buffers before dispatch."
731 ))
732 })?;
733 abi_input_bytes = abi_input_bytes.checked_add(input_len).ok_or_else(|| {
734 BackendError::new(
735 "megakernel ABI input byte total overflowed u64. Fix: shard ABI input buffers before dispatch.",
736 )
737 })?;
738 }
739 launch
740 .estimated_peak_device_bytes
741 .checked_add(abi_input_bytes)
742 .and_then(|bytes| bytes.checked_add(abi_input_bytes))
743 .ok_or_else(|| {
744 BackendError::new(
745 "megakernel peak device byte estimate overflowed u64. Fix: shard ABI buffers before dispatch.",
746 )
747 })
748}
749
750fn enforce_megakernel_device_memory_budget(
751 requested: u64,
752 available: u64,
753) -> Result<(), BackendError> {
754 if available != 0 && requested > available {
755 Err(BackendError::DeviceOutOfMemory {
756 requested,
757 available,
758 })
759 } else {
760 Ok(())
761 }
762}
763
764fn occupancy_proxy_bps(
765 item_count: usize,
766 worker_groups: u32,
767 workgroup_size_x: u32,
768) -> Result<u16, BackendError> {
769 let lanes = u64::from(worker_groups.max(1))
770 .checked_mul(u64::from(workgroup_size_x.max(1)))
771 .ok_or_else(|| {
772 BackendError::new(
773 "megakernel occupancy lane count overflowed u64. Fix: reduce worker groups or workgroup size.",
774 )
775 })?;
776 Ok(density_bps(
777 WGPU_NUMERIC.usize_to_u64(item_count, "megakernel occupancy item count")?,
778 lanes,
779 ))
780}
781
782fn density_bps(numerator: u64, denominator: u64) -> u16 {
783 crate::numeric::WGPU_NUMERIC
784 .ratio_basis_points_u64_wide(
785 numerator,
786 denominator.max(1),
787 0,
788 "megakernel occupancy density",
789 )
790 .min(10_000) as u16
791}
792
793fn ensure_empty_io_queue_bytes(bytes: &mut Vec<u8>) -> Result<(), BackendError> {
794 let expected = usize::try_from(IO_SLOT_COUNT)
795 .map_err(|source| {
796 BackendError::new(format!(
797 "IO_SLOT_COUNT cannot fit usize: {source}. Fix: keep IO_SLOT_COUNT within the host index ABI."
798 ))
799 })?
800 .checked_mul(usize::try_from(IO_SLOT_WORDS).map_err(|source| {
801 BackendError::new(format!(
802 "IO_SLOT_WORDS cannot fit usize: {source}. Fix: keep IO_SLOT_WORDS within the host index ABI."
803 ))
804 })?)
805 .and_then(|words| words.checked_mul(std::mem::size_of::<u32>()))
806 .ok_or_else(|| {
807 BackendError::new(
808 "megakernel IO queue byte length overflowed usize. Fix: shard IO queue slots before dispatch.".to_string(),
809 )
810 })?;
811 if bytes.len() != expected {
812 try_encode_empty_io_queue_into(IO_SLOT_COUNT, bytes)
813 .map_err(|error| BackendError::new(error.to_string()))?;
814 }
815 Ok(())
816}
817
818fn ensure_control_bytes(bytes: &mut Vec<u8>) -> Result<(), BackendError> {
819 let expected = Megakernel::control_byte_len(0).ok_or_else(|| {
820 BackendError::new(
821 "megakernel control byte length overflowed usize. Fix: reduce observable slot count."
822 .to_string(),
823 )
824 })?;
825 if bytes.len() != expected {
826 Megakernel::try_encode_control_into(false, 1, 0, bytes)
827 .map_err(|error| BackendError::new(error.to_string()))?;
828 }
829 Ok(())
830}
831
832fn ensure_empty_debug_log_bytes(bytes: &mut Vec<u8>) -> Result<(), BackendError> {
833 let record_capacity = Megakernel::debug_record_capacity();
834 let expected = Megakernel::debug_log_byte_len(record_capacity).ok_or_else(|| {
835 BackendError::new(
836 "megakernel debug-log byte length overflowed usize. Fix: reduce debug record capacity."
837 .to_string(),
838 )
839 })?;
840 if bytes.len() != expected {
841 Megakernel::try_encode_empty_debug_log_into(record_capacity, bytes)
842 .map_err(|error| BackendError::new(error.to_string()))?;
843 }
844 Ok(())
845}
846
847fn compiled_pipeline_cache_matches(
848 backend: &dyn VyreBackend,
849 workgroup_size_x: u32,
850 slot_count: u32,
851 dispatch_config: &DispatchConfig,
852 cache: &Option<CompiledMegakernelPipeline>,
853) -> bool {
854 cache.as_ref().is_some_and(|cached| {
855 cached.backend_id == backend.id()
856 && cached.backend_version == backend.version()
857 && cached.workgroup_size_x == workgroup_size_x
858 && cached.slot_count == slot_count
859 && same_dispatch_shape(&cached.dispatch_config, dispatch_config)
860 })
861}
862
863fn compiled_pipeline_for_geometry<'a>(
864 backend: &dyn VyreBackend,
865 program: Arc<Program>,
866 workgroup_size_x: u32,
867 slot_count: u32,
868 dispatch_config: &DispatchConfig,
869 cache: &'a mut Option<CompiledMegakernelPipeline>,
870) -> Result<Option<&'a dyn CompiledPipeline>, BackendError> {
871 if compiled_pipeline_cache_matches(
872 backend,
873 workgroup_size_x,
874 slot_count,
875 dispatch_config,
876 cache,
877 ) && cache
878 .as_ref()
879 .is_some_and(|cached| Arc::ptr_eq(&cached.program, &program))
880 {
881 return Ok(cache.as_ref().map(|cached| cached.pipeline.as_ref()));
882 }
883
884 match backend.compile_native_shared(program.clone(), dispatch_config)? {
885 Some(pipeline) => {
886 *cache = Some(CompiledMegakernelPipeline {
887 backend_id: backend.id(),
888 backend_version: backend.version(),
889 workgroup_size_x,
890 slot_count,
891 dispatch_config: dispatch_config.clone(),
892 program,
893 pipeline,
894 });
895 Ok(cache.as_ref().map(|cached| cached.pipeline.as_ref()))
896 }
897 None => {
898 *cache = None;
899 Ok(None)
900 }
901 }
902}
903
904fn same_dispatch_shape(left: &DispatchConfig, right: &DispatchConfig) -> bool {
905 left.profile == right.profile
906 && left.ulp_budget == right.ulp_budget
907 && left.max_output_bytes == right.max_output_bytes
908 && left.workgroup_override == right.workgroup_override
909 && left.grid_override == right.grid_override
910 && left.fixpoint_iterations == right.fixpoint_iterations
911 && left.speculation == right.speculation
912 && left.persistent_thread == right.persistent_thread
913 && left.cooperative == right.cooperative
914 && left.timeout == right.timeout
915}
916
917fn retained_redundant_done_count(
918 work_items: &[MegakernelWorkItem],
919 dispatch_items: &[MegakernelWorkItem],
920 done_count: u64,
921 dispatch_item_count: usize,
922 redundancy: &CrossArmRedundancy,
923) -> u64 {
924 let dispatch_item_count = u64::try_from(dispatch_item_count).unwrap_or(u64::MAX);
925 if done_count < dispatch_item_count {
926 return 0;
927 }
928 redundancy
929 .redundant_pairs
930 .iter()
931 .filter(|(early_idx, _, _)| {
932 work_items
933 .get(*early_idx)
934 .is_some_and(|item| dispatch_items.iter().any(|queued| queued == item))
935 })
936 .count()
937 .try_into()
938 .unwrap_or(u64::MAX)
939}
940
941fn nanos_u64(nanos: u128) -> Result<u64, BackendError> {
942 u64::try_from(nanos).map_err(|source| {
943 BackendError::new(format!(
944 "megakernel elapsed time cannot fit u64 nanoseconds: {source}. Fix: split or timeout the dispatch before telemetry overflows."
945 ))
946 })
947}
948
949fn should_track_lineage(item_count: usize) -> bool {
950 item_count <= MAX_INLINE_LINEAGE_ITEMS
951}
952
953impl MegakernelDispatch for WgpuMegakernelDispatcher<'_> {
954 fn dispatch_megakernel(
955 &self,
956 work_queue: &[MegakernelWorkItem],
957 config: &MegakernelConfig,
958 ) -> Result<MegakernelReport, BackendError> {
959 WgpuMegakernelDispatcher::dispatch_megakernel(self, work_queue, config)
960 }
961}
962
963#[cfg(test)]
964mod tests {
965 use super::*;
966
967 fn item(op: u32, input: u32, output: u32, param: u32) -> MegakernelWorkItem {
968 MegakernelWorkItem {
969 op_handle: op,
970 input_handle: input,
971 output_handle: output,
972 param,
973 }
974 }
975
976 #[test]
977 fn retained_redundant_done_count_is_zero_without_full_dispatch_completion() {
978 let a = item(1, 0, 5, 7);
979 let work_items = [a, a];
980 let redundancy = CrossArmRedundancy {
981 redundant_pairs: vec![(0, 1, 0)],
982 total_redundant_ops: 1,
983 };
984
985 let count = retained_redundant_done_count(&work_items, &[a], 0, 1, &redundancy);
986
987 assert_eq!(count, 0);
988 }
989
990 #[test]
991 fn retained_redundant_done_count_counts_duplicates_when_producer_finished() {
992 let a = item(1, 0, 5, 7);
993 let b = item(2, 5, 6, 0);
994 let work_items = [a, b, a, a];
995 let redundancy = CrossArmRedundancy {
996 redundant_pairs: vec![(0, 2, 0), (0, 3, 0)],
997 total_redundant_ops: 2,
998 };
999
1000 let count = retained_redundant_done_count(&work_items, &[a, b], 2, 2, &redundancy);
1001
1002 assert_eq!(count, 2);
1003 }
1004
1005 #[test]
1006 fn retained_redundant_done_count_ignores_redundancy_without_queued_producer() {
1007 let a = item(1, 0, 5, 7);
1008 let b = item(2, 5, 6, 0);
1009 let work_items = [a, a, b];
1010 let redundancy = CrossArmRedundancy {
1011 redundant_pairs: vec![(0, 1, 0)],
1012 total_redundant_ops: 1,
1013 };
1014
1015 let count = retained_redundant_done_count(&work_items, &[b], 1, 1, &redundancy);
1016
1017 assert_eq!(count, 0);
1018 }
1019
1020 #[test]
1021 fn retained_redundant_done_count_ignores_invalid_indices() {
1022 let a = item(1, 0, 5, 7);
1023 let redundancy = CrossArmRedundancy {
1024 redundant_pairs: vec![(99, 1, 0)],
1025 total_redundant_ops: 1,
1026 };
1027
1028 let count = retained_redundant_done_count(&[a], &[a], 1, 1, &redundancy);
1029
1030 assert_eq!(count, 0);
1031 }
1032
1033 #[test]
1034 fn lineage_tracking_is_capped_for_large_hot_queues() {
1035 assert!(should_track_lineage(MAX_INLINE_LINEAGE_ITEMS));
1036 assert!(!should_track_lineage(MAX_INLINE_LINEAGE_ITEMS + 1));
1037 }
1038
1039 #[test]
1040 fn dispatch_shape_distinguishes_timeout() {
1041 let mut left = DispatchConfig::default();
1042 let mut right = DispatchConfig::default();
1043 left.timeout = Some(std::time::Duration::from_millis(1));
1044 right.timeout = Some(std::time::Duration::from_millis(2));
1045
1046 assert!(!same_dispatch_shape(&left, &right));
1047 }
1048
1049 #[test]
1050 fn megakernel_dispatch_uses_caller_owned_output_scratch() {
1051 let source = include_str!("megakernel.rs");
1052
1053 assert!(
1054 !source.contains(concat!("scratch.outputs", ".clear();")),
1055 "Fix: megakernel dispatch must not drop reusable output slots before dispatch."
1056 );
1057 assert!(
1058 !source.contains(concat!("scratch.outputs", " =")),
1059 "Fix: megakernel dispatch must not replace caller-owned output scratch with fresh OutputBuffers."
1060 );
1061 assert!(
1062 source.contains("dispatch_persistent_handles_into("),
1063 "Fix: resident megakernel dispatch must collect into reusable scratch outputs."
1064 );
1065 assert!(
1066 source.contains("dispatch_borrowed_into("),
1067 "Fix: borrowed megakernel dispatch must collect into reusable scratch outputs."
1068 );
1069 }
1070
1071 #[test]
1072 fn dispatch_rejects_device_memory_budget_before_backend_work() {
1073 let backend = FakeCudaResidentBackend::new();
1074 let dispatcher = WgpuMegakernelDispatcher::new(&backend);
1075 let config = MegakernelConfig {
1076 worker_count: 1,
1077 workload: vyre_runtime::megakernel::MegakernelWorkloadHints {
1078 resident_device_bytes: 128 * 1024,
1079 device_memory_budget_bytes: 64 * 1024,
1080 ..Default::default()
1081 },
1082 ..MegakernelConfig::default()
1083 };
1084
1085 let error = dispatcher
1086 .dispatch_megakernel(&[item(1, 0, 1, 0)], &config)
1087 .expect_err("over-budget megakernel launch must fail before backend work");
1088
1089 match error {
1090 BackendError::DeviceOutOfMemory {
1091 requested,
1092 available,
1093 } => {
1094 assert!(
1095 requested > available,
1096 "Fix: megakernel budget rejection must report requested bytes above the budget."
1097 );
1098 assert_eq!(available, 64 * 1024);
1099 }
1100 other => panic!(
1101 "expected structured DeviceOutOfMemory for over-budget launch, got {other:?}"
1102 ),
1103 }
1104 assert!(
1105 backend.uploads.lock().unwrap().is_empty(),
1106 "Fix: over-budget megakernel launch must fail before resident uploads."
1107 );
1108 assert!(
1109 backend.frees.lock().unwrap().is_empty(),
1110 "Fix: over-budget megakernel launch must not allocate resources that need cleanup."
1111 );
1112 }
1113
1114 #[test]
1115 fn dispatch_rejects_abi_buffer_budget_before_backend_work() {
1116 let backend = FakeCudaResidentBackend::new();
1117 let dispatcher = WgpuMegakernelDispatcher::new(&backend);
1118 let launch = MegakernelConfig::default()
1119 .launch_recommendation(1, 1, 1, 1)
1120 .expect("Fix: test launch recommendation must be valid");
1121 let budget_above_policy_scratch = launch
1122 .estimated_peak_device_bytes
1123 .checked_add(1)
1124 .expect("Fix: WGPU megakernel peak-device-byte policy sentinel overflowed u64");
1125 let config = MegakernelConfig {
1126 worker_count: 1,
1127 workload: vyre_runtime::megakernel::MegakernelWorkloadHints {
1128 device_memory_budget_bytes: budget_above_policy_scratch,
1129 ..Default::default()
1130 },
1131 ..MegakernelConfig::default()
1132 };
1133
1134 let error = dispatcher
1135 .dispatch_megakernel(&[item(1, 0, 1, 0)], &config)
1136 .expect_err("ABI buffers must be included in megakernel device budget");
1137
1138 match error {
1139 BackendError::DeviceOutOfMemory {
1140 requested,
1141 available,
1142 } => {
1143 assert!(
1144 requested > available,
1145 "Fix: megakernel ABI budget rejection must report actual peak bytes above the budget."
1146 );
1147 assert_eq!(available, budget_above_policy_scratch);
1148 }
1149 other => panic!(
1150 "expected structured DeviceOutOfMemory for ABI-buffer budget overflow, got {other:?}"
1151 ),
1152 }
1153 assert!(
1154 backend.uploads.lock().unwrap().is_empty(),
1155 "Fix: ABI-budget rejection must happen before resident uploads."
1156 );
1157 }
1158
1159 #[test]
1160 fn resident_input_upload_plan_covers_every_abi_slot_in_order() {
1161 let owner = vyre_driver::ResidentOwner::new()
1162 .expect("Fix: resident owner minting must succeed in tests");
1163 let resources = vec![
1164 Resource::Resident(owner.handle(10)),
1165 Resource::Resident(owner.handle(11)),
1166 Resource::Resident(owner.handle(12)),
1167 Resource::Resident(owner.handle(13)),
1168 ];
1169 let inputs: [&[u8]; 4] = [
1170 &b"control"[..],
1171 &b"ring"[..],
1172 &b"debug"[..],
1173 &b"io_queue"[..],
1174 ];
1175
1176 let plan = resident_input_upload_plan(&resources, &inputs)
1177 .expect("Fix: resident ABI upload plan should cover exact megakernel input slots");
1178
1179 assert_eq!(plan.len(), inputs.len());
1180 for index in 0..inputs.len() {
1181 assert!(
1182 std::ptr::eq(plan[index].0, &resources[index]),
1183 "Fix: resident megakernel upload slot {index} must target the matching resource."
1184 );
1185 assert_eq!(
1186 plan[index].1, inputs[index],
1187 "Fix: resident megakernel upload slot {index} must refresh the matching input bytes."
1188 );
1189 }
1190 }
1191
1192 #[test]
1193 fn resident_input_upload_plan_rejects_abi_slot_mismatch() {
1194 let owner = vyre_driver::ResidentOwner::new()
1195 .expect("Fix: resident owner minting must succeed in tests");
1196 let resources = vec![
1197 Resource::Resident(owner.handle(10)),
1198 Resource::Resident(owner.handle(11)),
1199 ];
1200 let inputs: [&[u8]; 4] = [&b"control"[..], &b"ring"[..], &b"debug"[..], &b"io"[..]];
1201
1202 let error = resident_input_upload_plan(&resources, &inputs)
1203 .expect_err("resident upload plan must reject truncated resource lists");
1204
1205 assert!(
1206 error
1207 .to_string()
1208 .contains("expected 4 resident slot(s) for 2 input buffer(s)"),
1209 "Fix: resident ABI slot mismatch diagnostics must include expected and actual counts."
1210 );
1211 }
1212
1213 #[test]
1214 fn megakernel_report_telemetry_counts_transfer_and_density() {
1215 let inputs: [&[u8]; 4] = [&[0; 4][..], &[1; 8][..], &[2; 12][..], &[3; 16][..]];
1216 let outputs = vec![vec![0; 5], vec![0; 7], vec![0; 11], vec![0; 13]];
1217
1218 let launch = vyre_runtime::megakernel::MegakernelLaunchPolicy::standard()
1219 .recommend(vyre_runtime::megakernel::MegakernelLaunchRequest::direct(
1220 8, 2, 8,
1221 ))
1222 .expect("Fix: test launch policy request must be valid");
1223 let peak_bytes = megakernel_dispatch_peak_device_bytes(&inputs, launch)
1224 .expect("Fix: test peak byte estimate must fit u64");
1225 let telemetry = megakernel_report_telemetry(
1226 &inputs, &outputs, 4, 8, 16, 2, 8, launch, peak_bytes, true, false,
1227 )
1228 .expect("Fix: test telemetry byte accounting must fit u64");
1229
1230 assert_eq!(telemetry.bytes_uploaded, 40);
1231 assert_eq!(telemetry.bytes_read_back, 36);
1232 assert_eq!(telemetry.bytes_moved, 76);
1233 assert_eq!(telemetry.resident_allocations, 4);
1234 assert_eq!(telemetry.kernel_launches, 1);
1235 assert_eq!(telemetry.sync_points, 1);
1236 assert_eq!(telemetry.occupancy_proxy_bps, 5_000);
1237 assert_eq!(telemetry.frontier_density_bps, 5_000);
1238 assert_eq!(density_bps(u64::MAX, 1), 10_000);
1239 assert_eq!(telemetry.readback_buffers, 4);
1240 assert!(telemetry.compiled_pipeline_cache_hit);
1241 assert!(!telemetry.resident_input_cache_hit);
1242 assert_eq!(telemetry.topology, launch.topology);
1243 assert_eq!(telemetry.pressure, launch.pressure);
1244 assert_eq!(telemetry.execution_mode, launch.execution_mode);
1245 assert_eq!(telemetry.hit_capacity, launch.hit_capacity);
1246 assert_eq!(telemetry.estimated_peak_device_bytes, peak_bytes);
1247 assert_eq!(
1248 telemetry.device_memory_budget_bytes,
1249 launch.device_memory_budget_bytes
1250 );
1251 }
1252
1253 struct FakeCudaResidentBackend {
1254 version: &'static str,
1255 owner: vyre_driver::ResidentOwner,
1256 next: std::sync::atomic::AtomicU64,
1257 uploads: std::sync::Mutex<Vec<(vyre_driver::ResidentHandle, usize)>>,
1258 frees: std::sync::Mutex<Vec<vyre_driver::ResidentHandle>>,
1259 supported_ops: std::collections::HashSet<vyre_foundation::ir::OpId>,
1260 }
1261
1262 impl FakeCudaResidentBackend {
1263 fn new() -> Self {
1264 Self::with_version("test-v1")
1265 }
1266
1267 fn with_version(version: &'static str) -> Self {
1268 Self {
1269 version,
1270 owner: vyre_driver::ResidentOwner::new()
1271 .expect("Fix: resident owner minting must succeed in tests"),
1272 next: std::sync::atomic::AtomicU64::new(100),
1273 uploads: std::sync::Mutex::new(Vec::new()),
1274 frees: std::sync::Mutex::new(Vec::new()),
1275 supported_ops: std::collections::HashSet::new(),
1276 }
1277 }
1278 }
1279
1280 impl vyre_driver::backend::private::Sealed for FakeCudaResidentBackend {}
1281
1282 impl VyreBackend for FakeCudaResidentBackend {
1283 fn id(&self) -> &'static str {
1284 "cuda"
1285 }
1286
1287 fn version(&self) -> &'static str {
1288 self.version
1289 }
1290
1291 fn supported_ops(&self) -> &std::collections::HashSet<vyre_foundation::ir::OpId> {
1292 &self.supported_ops
1293 }
1294
1295 fn dispatch(
1296 &self,
1297 _program: &Program,
1298 _inputs: &[Vec<u8>],
1299 _config: &DispatchConfig,
1300 ) -> Result<Vec<Vec<u8>>, BackendError> {
1301 Err(BackendError::new(
1302 "fake CUDA resident backend must not run host dispatch in resident-cache tests.",
1303 ))
1304 }
1305
1306 fn allocate_resident(&self, _byte_len: usize) -> Result<Resource, BackendError> {
1307 Ok(Resource::Resident(self.owner.handle(
1308 self.next.fetch_add(1, std::sync::atomic::Ordering::Relaxed),
1309 )))
1310 }
1311
1312 fn upload_resident_many(&self, uploads: &[(&Resource, &[u8])]) -> Result<(), BackendError> {
1313 let mut captured = self.uploads.lock().map_err(BackendError::poisoned_lock)?;
1314 for &(resource, bytes) in uploads {
1315 let Resource::Resident(handle) = resource else {
1316 return Err(BackendError::new(
1317 "fake CUDA resident backend expected resident handles.",
1318 ));
1319 };
1320 captured.push((*handle, bytes.len()));
1321 }
1322 Ok(())
1323 }
1324
1325 fn free_resident(&self, resource: Resource) -> Result<(), BackendError> {
1326 let Resource::Resident(handle) = resource else {
1327 return Err(BackendError::new(
1328 "fake CUDA resident backend expected resident handles for free.",
1329 ));
1330 };
1331 self.frees
1332 .lock()
1333 .map_err(BackendError::poisoned_lock)?
1334 .push(handle);
1335 Ok(())
1336 }
1337 }
1338
1339 struct FakeCompiledPipeline;
1340
1341 impl vyre_driver::backend::private::Sealed for FakeCompiledPipeline {}
1342
1343 impl CompiledPipeline for FakeCompiledPipeline {
1344 fn id(&self) -> &str {
1345 "fake-compiled-megakernel-pipeline"
1346 }
1347
1348 fn dispatch(
1349 &self,
1350 _inputs: &[Vec<u8>],
1351 _config: &DispatchConfig,
1352 ) -> Result<Vec<Vec<u8>>, BackendError> {
1353 Err(BackendError::new(
1354 "fake compiled pipeline must not dispatch in cache identity tests.",
1355 ))
1356 }
1357 }
1358
1359 struct FakeCudaNoResidentBackend {
1360 supported_ops: std::collections::HashSet<vyre_foundation::ir::OpId>,
1361 }
1362
1363 impl FakeCudaNoResidentBackend {
1364 fn new() -> Self {
1365 Self {
1366 supported_ops: std::collections::HashSet::new(),
1367 }
1368 }
1369 }
1370
1371 impl vyre_driver::backend::private::Sealed for FakeCudaNoResidentBackend {}
1372
1373 impl VyreBackend for FakeCudaNoResidentBackend {
1374 fn id(&self) -> &'static str {
1375 "cuda"
1376 }
1377
1378 fn supported_ops(&self) -> &std::collections::HashSet<vyre_foundation::ir::OpId> {
1379 &self.supported_ops
1380 }
1381
1382 fn dispatch(
1383 &self,
1384 _program: &Program,
1385 _inputs: &[Vec<u8>],
1386 _config: &DispatchConfig,
1387 ) -> Result<Vec<Vec<u8>>, BackendError> {
1388 Err(BackendError::new(
1389 "fake no-resident CUDA backend must not fall back to host dispatch.",
1390 ))
1391 }
1392
1393 fn allocate_resident(&self, _byte_len: usize) -> Result<Resource, BackendError> {
1394 Err(BackendError::UnsupportedFeature {
1395 name: "resident allocation".to_string(),
1396 backend: self.id().to_string(),
1397 })
1398 }
1399 }
1400
1401 #[test]
1402 fn cuda_resident_megakernel_allocation_failure_is_not_borrowed_fallback() {
1403 let backend = FakeCudaNoResidentBackend::new();
1404 let inputs: [&[u8]; 4] = [&[1; 4][..], &[2; 8][..], &[3; 12][..], &[4; 16][..]];
1405 let mut cache = None;
1406
1407 let error = ensure_resident_megakernel_buffers(&backend, 64, 64, &inputs, &mut cache)
1408 .expect_err(
1409 "CUDA megakernel resident setup must fail loudly when residency is unsupported",
1410 );
1411
1412 assert!(
1413 error
1414 .to_string()
1415 .contains("CUDA resident megakernel input allocation"),
1416 "Fix: CUDA megakernel residency failure must not be hidden behind borrowed dispatch fallback: {error}"
1417 );
1418 assert!(
1419 cache.is_none(),
1420 "Fix: failed CUDA resident setup must not leave a partial cache entry."
1421 );
1422 }
1423
1424 #[test]
1425 fn compiled_megakernel_cache_separates_backend_versions() {
1426 let first_backend = FakeCudaResidentBackend::with_version("test-v1");
1427 let second_backend = FakeCudaResidentBackend::with_version("test-v2");
1428 let dispatch_config = DispatchConfig::default();
1429 let cache = Some(CompiledMegakernelPipeline {
1430 backend_id: first_backend.id(),
1431 backend_version: first_backend.version(),
1432 workgroup_size_x: 64,
1433 slot_count: 64,
1434 dispatch_config: dispatch_config.clone(),
1435 program: Arc::new(Program::default()),
1436 pipeline: Arc::new(FakeCompiledPipeline),
1437 });
1438
1439 assert!(
1440 compiled_pipeline_cache_matches(&first_backend, 64, 64, &dispatch_config, &cache),
1441 "Fix: identical backend version and megakernel geometry must allow compiled cache reuse."
1442 );
1443 assert!(
1444 !compiled_pipeline_cache_matches(&second_backend, 64, 64, &dispatch_config, &cache),
1445 "Fix: compiled megakernel cache identity must include backend implementation version."
1446 );
1447 }
1448
1449 #[test]
1450 fn resident_megakernel_buffers_reuse_resources_and_refresh_all_inputs() {
1451 let backend = FakeCudaResidentBackend::new();
1452 let inputs: [&[u8]; 4] = [&[1; 4][..], &[2; 8][..], &[3; 12][..], &[4; 16][..]];
1453 let mut cache = None;
1454
1455 let first_resources =
1456 ensure_resident_megakernel_buffers(&backend, 64, 64, &inputs, &mut cache)
1457 .expect("Fix: first resident ensure must succeed")
1458 .expect("Fix: fake CUDA backend supports resident resources")
1459 .to_vec();
1460 let first_uploads = backend.uploads.lock().unwrap().clone();
1461 let second_resources =
1462 ensure_resident_megakernel_buffers(&backend, 64, 64, &inputs, &mut cache)
1463 .expect("Fix: second resident ensure must succeed")
1464 .expect("Fix: fake CUDA backend supports resident resources")
1465 .to_vec();
1466 let all_uploads = backend.uploads.lock().unwrap().clone();
1467
1468 assert_eq!(
1469 first_resources, second_resources,
1470 "Fix: identical megakernel resident input shapes must reuse resident resources."
1471 );
1472 let slot = |id: u64, byte_len: usize| (backend.owner.handle(id), byte_len);
1473 assert_eq!(
1474 first_uploads,
1475 vec![slot(100, 4), slot(101, 8), slot(102, 12), slot(103, 16)],
1476 "Fix: first resident publication must upload every megakernel ABI input slot."
1477 );
1478 assert_eq!(
1479 all_uploads,
1480 vec![
1481 slot(100, 4),
1482 slot(101, 8),
1483 slot(102, 12),
1484 slot(103, 16),
1485 slot(100, 4),
1486 slot(101, 8),
1487 slot(102, 12),
1488 slot(103, 16)
1489 ],
1490 "Fix: cache-hit resident publication must refresh all four volatile ABI input slots."
1491 );
1492 assert!(
1493 backend.frees.lock().unwrap().is_empty(),
1494 "Fix: cache-hit resident publication must not free and reallocate resources."
1495 );
1496 }
1497
1498 #[test]
1499 fn resident_megakernel_cache_separates_backend_versions() {
1500 let first_backend = FakeCudaResidentBackend::with_version("test-v1");
1501 let second_backend = FakeCudaResidentBackend::with_version("test-v2");
1502 let inputs: [&[u8]; 4] = [&[1; 4][..], &[2; 8][..], &[3; 12][..], &[4; 16][..]];
1503 let mut cache = None;
1504
1505 ensure_resident_megakernel_buffers(&first_backend, 64, 64, &inputs, &mut cache)
1506 .expect("Fix: first resident ensure must succeed")
1507 .expect("Fix: first fake CUDA backend supports resident resources");
1508
1509 assert!(
1510 !resident_megakernel_cache_matches(
1511 &second_backend,
1512 64,
1513 64,
1514 [4, 8, 12, 16],
1515 &cache
1516 ),
1517 "Fix: megakernel resident input cache identity must include backend implementation version."
1518 );
1519 }
1520}