Skip to main content

vyre_driver_wgpu/
megakernel.rs

1//! WGPU-owned megakernel dispatch wrapper.
2
3use 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
105/// Runtime wrapper for persistent megakernel dispatch.
106pub struct WgpuMegakernelDispatcher<'a> {
107    backend: &'a dyn VyreBackend,
108}
109
110impl<'a> WgpuMegakernelDispatcher<'a> {
111    /// Create a new dispatcher.
112    #[must_use]
113    pub fn new(backend: &'a dyn VyreBackend) -> Self {
114        Self { backend }
115    }
116
117    /// Decode a raw little-endian `MegakernelWorkItem` queue and launch the megakernel.
118    ///
119    /// # Errors
120    ///
121    /// Returns a backend error when `work_queue_bytes` is not exactly aligned to
122    /// [`MegakernelWorkItem`] records or when backend dispatch fails.
123    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    /// Launch the megakernel.
144    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    /// Launch the megakernel with a caller-supplied IO queue.
168    ///
169    /// The queue is validated against the megakernel ABI before any backend
170    /// work starts, so malformed queue views fail before compilation or GPU
171    /// submission.
172    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        // P-RUNTIME-1: attach scallop-provenance lineage per dispatched
383        // region so observability collectors can attribute outputs back
384        // to the source rules that derived them. We seed the lineage
385        // bitset from work_items[i].op_handle (each op contributes its
386        // own bit, capped at 32 distinct ops per dispatch  -  the u32
387        // word width) and run the substrate provenance closure across
388        // the same exchange_adj that the matroid scheduler used. The
389        // closure propagates lineage through any fused-region edges,
390        // so a fused region's lineage bitset = union of contributing
391        // ops' bits.
392        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}