1use crate::{megakernel::planner::MegakernelWorkItem, PipelineError};
21use rustc_hash::FxHashMap;
22use vyre_foundation::allocation::{try_reserve_hash_map_to_capacity, try_reserve_vec_to_capacity};
23
24const DENSE_OUTPUT_UNIQUE_BITS: usize = 4096;
25const DENSE_OUTPUT_UNIQUE_WORDS: usize = DENSE_OUTPUT_UNIQUE_BITS / u64::BITS as usize;
26
27#[derive(Debug, Clone, Default, PartialEq, Eq)]
35pub struct CrossArmRedundancy {
36 pub redundant_pairs: Vec<(usize, usize, usize)>,
40 pub total_redundant_ops: usize,
44}
45
46impl CrossArmRedundancy {
47 #[must_use]
49 pub fn new() -> Self {
50 Self::default()
51 }
52
53 #[must_use]
55 pub fn is_empty(&self) -> bool {
56 self.redundant_pairs.is_empty()
57 }
58}
59
60#[derive(Debug, Default)]
62pub struct RedundantWorkItemPruneScratch {
63 first_seen: FxHashMap<(u32, u32, u32, u32), usize>,
64}
65
66impl RedundantWorkItemPruneScratch {
67 pub fn clear(&mut self) {
69 self.first_seen.clear();
70 }
71
72 fn try_prepare_for_len(&mut self, len: usize) -> Result<(), PipelineError> {
73 self.first_seen.clear();
74 let retained_ceiling = len.checked_mul(4).unwrap_or(usize::MAX).max(1024);
75 if self.first_seen.capacity() > retained_ceiling {
76 self.first_seen.shrink_to(len);
77 }
78 if self.first_seen.capacity() < len {
79 try_reserve_hash_map_to_capacity(&mut self.first_seen, len).map_err(|source| {
80 PipelineError::Backend(format!(
81 "megakernel redundant-work hash reservation failed for {len} item(s): {source}. Fix: shard the work batch before pruning."
82 ))
83 })?;
84 }
85 Ok(())
86 }
87}
88
89#[must_use]
101#[cfg(any(test, feature = "legacy-infallible"))]
102pub fn detect_cross_arm_redundancy(arms: &[&[MegakernelWorkItem]]) -> CrossArmRedundancy {
103 try_detect_cross_arm_redundancy(arms).unwrap_or_else(|error| {
104 panic!(
105 "megakernel cross-arm redundancy detection allocation failed: {error}. Fix: split the fused arm sequence before planning."
106 )
107 })
108}
109
110pub fn try_detect_cross_arm_redundancy(
117 arms: &[&[MegakernelWorkItem]],
118) -> Result<CrossArmRedundancy, PipelineError> {
119 let total_ops = arms.iter().map(|arm| arm.len()).sum();
121 let mut first_seen: FxHashMap<(u32, u32, u32), usize> = FxHashMap::default();
122 reserve_hash_map(&mut first_seen, total_ops, "cross-arm first-seen")?;
123 let mut report = CrossArmRedundancy {
124 redundant_pairs: Vec::new(),
125 total_redundant_ops: 0,
126 };
127 for (arm_idx, arm) in arms.iter().enumerate() {
128 for (op_idx, item) in arm.iter().enumerate() {
129 let key = (item.op_handle, item.input_handle, item.output_handle);
130 match first_seen.get(&key) {
131 Some(&early_arm_idx) if early_arm_idx < arm_idx => {
132 reserve_redundant_pairs(&mut report.redundant_pairs, 1, "cross-arm report")?;
133 report
134 .redundant_pairs
135 .push((early_arm_idx, arm_idx, op_idx));
136 }
137 Some(_) => {
138 }
140 None => {
141 first_seen.insert(key, arm_idx);
142 }
143 }
144 }
145 }
146 report.total_redundant_ops = report.redundant_pairs.len();
147 Ok(report)
148}
149
150#[cfg(any(test, feature = "legacy-infallible"))]
166pub fn prune_redundant_work_items_into(
167 items: &[MegakernelWorkItem],
168 out: &mut Vec<MegakernelWorkItem>,
169) -> CrossArmRedundancy {
170 try_prune_redundant_work_items_into(items, out).unwrap_or_else(|error| {
171 panic!(
172 "megakernel redundant-work pruning allocation failed: {error}. Fix: shard the work batch before pruning."
173 )
174 })
175}
176
177pub fn try_prune_redundant_work_items_into(
185 items: &[MegakernelWorkItem],
186 out: &mut Vec<MegakernelWorkItem>,
187) -> Result<CrossArmRedundancy, PipelineError> {
188 let mut scratch = RedundantWorkItemPruneScratch::default();
189 try_prune_redundant_work_items_with_scratch_into(items, out, &mut scratch)
190}
191
192#[cfg(any(test, feature = "legacy-infallible"))]
199pub fn prune_redundant_work_items_with_scratch_into(
200 items: &[MegakernelWorkItem],
201 out: &mut Vec<MegakernelWorkItem>,
202 scratch: &mut RedundantWorkItemPruneScratch,
203) -> CrossArmRedundancy {
204 try_prune_redundant_work_items_with_scratch_into(items, out, scratch).unwrap_or_else(|error| {
205 panic!(
206 "megakernel redundant-work pruning allocation failed: {error}. Fix: shard the work batch before pruning."
207 )
208 })
209}
210
211pub fn try_prune_redundant_work_items_with_scratch_into(
219 items: &[MegakernelWorkItem],
220 out: &mut Vec<MegakernelWorkItem>,
221 scratch: &mut RedundantWorkItemPruneScratch,
222) -> Result<CrossArmRedundancy, PipelineError> {
223 out.clear();
224
225 if output_handles_are_dense_unique(items) {
226 scratch.clear();
227 return Ok(CrossArmRedundancy::new());
228 }
229
230 scratch.try_prepare_for_len(items.len())?;
231 let mut report = CrossArmRedundancy {
232 redundant_pairs: Vec::new(),
233 total_redundant_ops: 0,
234 };
235 let mut found_duplicate = false;
236
237 for (idx, item) in items.iter().copied().enumerate() {
238 let key = (
239 item.op_handle,
240 item.input_handle,
241 item.output_handle,
242 item.param,
243 );
244 if let Some(&early_idx) = scratch.first_seen.get(&key) {
245 if !found_duplicate {
246 reserve_work_items(out, items.len().checked_sub(1).unwrap_or(0), "dedup output")?;
247 out.extend_from_slice(&items[..idx]);
248 found_duplicate = true;
249 }
250 reserve_redundant_pairs(&mut report.redundant_pairs, 1, "dedup report")?;
251 report.redundant_pairs.push((early_idx, idx, 0));
252 continue;
253 }
254 scratch.first_seen.insert(key, idx);
255 if found_duplicate {
256 out.push(item);
257 }
258 }
259
260 report.total_redundant_ops = report.redundant_pairs.len();
261 Ok(report)
262}
263
264fn reserve_hash_map<K, V>(
265 values: &mut FxHashMap<K, V>,
266 additional: usize,
267 label: &'static str,
268) -> Result<(), PipelineError>
269where
270 K: Eq + std::hash::Hash,
271{
272 if additional > 0 {
273 let capacity = values.len().checked_add(additional).ok_or_else(|| {
274 PipelineError::Backend(format!(
275 "megakernel {label} reservation overflowed for {additional} additional entry slot(s). Fix: shard the work batch before whole-megakernel optimization."
276 ))
277 })?;
278 try_reserve_hash_map_to_capacity(values, capacity).map_err(|source| {
279 PipelineError::Backend(format!(
280 "megakernel {label} reservation failed for {additional} additional entry slot(s): {source}. Fix: shard the work batch before whole-megakernel optimization."
281 ))
282 })?;
283 }
284 Ok(())
285}
286
287fn reserve_redundant_pairs(
288 values: &mut Vec<(usize, usize, usize)>,
289 additional: usize,
290 label: &'static str,
291) -> Result<(), PipelineError> {
292 values.try_reserve(additional).map_err(|source| {
293 PipelineError::Backend(format!(
294 "megakernel {label} reservation failed for {additional} additional pair slot(s): {source}. Fix: shard the work batch before whole-megakernel optimization."
295 ))
296 })
297}
298
299fn reserve_work_items(
300 values: &mut Vec<MegakernelWorkItem>,
301 capacity: usize,
302 label: &'static str,
303) -> Result<(), PipelineError> {
304 if values.capacity() < capacity {
305 try_reserve_vec_to_capacity(values, capacity).map_err(|source| {
306 PipelineError::Backend(format!(
307 "megakernel {label} reservation failed for {capacity} item slot(s): {source}. Fix: shard the work batch before whole-megakernel optimization."
308 ))
309 })?;
310 }
311 Ok(())
312}
313
314fn output_handles_are_dense_unique(items: &[MegakernelWorkItem]) -> bool {
315 if items.len() <= 1 {
316 return true;
317 }
318 if items.len() > DENSE_OUTPUT_UNIQUE_BITS {
319 return false;
320 }
321
322 let mut min = u32::MAX;
323 let mut max = 0u32;
324 for item in items {
325 min = min.min(item.output_handle);
326 max = max.max(item.output_handle);
327 }
328 let Some(range) = u64::from(max)
329 .checked_sub(u64::from(min))
330 .and_then(|value| value.checked_add(1))
331 else {
332 return false;
333 };
334 if range > DENSE_OUTPUT_UNIQUE_BITS as u64 {
335 return false;
336 }
337
338 let mut seen = [0u64; DENSE_OUTPUT_UNIQUE_WORDS];
339 for item in items {
340 let Some(delta) = item.output_handle.checked_sub(min) else {
341 return false;
342 };
343 let Ok(offset) = usize::try_from(delta) else {
344 return false;
345 };
346 let word = offset / 64;
347 let bit = 1u64 << (offset % 64);
348 if (seen[word] & bit) != 0 {
349 return false;
350 }
351 seen[word] |= bit;
352 }
353 true
354}
355
356#[cfg(test)]
357mod tests {
358 use super::*;
359
360 fn item(op: u32, inp: u32, out: u32) -> MegakernelWorkItem {
361 MegakernelWorkItem {
362 op_handle: op,
363 input_handle: inp,
364 output_handle: out,
365 param: 0,
366 }
367 }
368
369 #[test]
370 fn empty_arms_have_no_redundancy() {
371 let arms: [&[MegakernelWorkItem]; 0] = [];
372 assert_eq!(
373 detect_cross_arm_redundancy(&arms),
374 CrossArmRedundancy::new()
375 );
376 }
377
378 #[test]
379 fn single_arm_with_repeats_has_no_cross_arm_redundancy() {
380 let a = vec![item(1, 0, 5), item(1, 0, 5), item(2, 5, 6)];
381 let arms: [&[MegakernelWorkItem]; 1] = [&a];
382 let report = detect_cross_arm_redundancy(&arms);
383 assert!(report.is_empty(), "intra-arm repeats are not cross-arm");
384 assert_eq!(report.total_redundant_ops, 0);
385 }
386
387 #[test]
388 fn identical_arms_report_full_overlap() {
389 let a = vec![item(1, 0, 5), item(2, 5, 6)];
390 let b = vec![item(1, 0, 5), item(2, 5, 6)];
391 let arms: [&[MegakernelWorkItem]; 2] = [&a, &b];
392 let report = detect_cross_arm_redundancy(&arms);
393 assert_eq!(report.total_redundant_ops, 2);
394 assert_eq!(report.redundant_pairs, vec![(0, 1, 0), (0, 1, 1)]);
395 }
396
397 #[test]
398 fn fully_disjoint_arms_have_no_redundancy() {
399 let a = vec![item(1, 0, 5)];
400 let b = vec![item(2, 7, 8)];
401 let arms: [&[MegakernelWorkItem]; 2] = [&a, &b];
402 assert!(detect_cross_arm_redundancy(&arms).is_empty());
403 }
404
405 #[test]
406 fn redundancy_uses_first_seen_arm_index() {
407 let a = vec![item(1, 0, 5)];
409 let b = vec![item(99, 0, 0)];
410 let c = vec![item(1, 0, 5)];
411 let d = vec![item(1, 0, 5)];
412 let arms: [&[MegakernelWorkItem]; 4] = [&a, &b, &c, &d];
413 let report = detect_cross_arm_redundancy(&arms);
414 assert_eq!(report.total_redundant_ops, 2);
415 assert_eq!(report.redundant_pairs, vec![(0, 2, 0), (0, 3, 0)]);
416 }
417
418 #[test]
419 fn param_field_does_not_affect_redundancy() {
420 let a = vec![MegakernelWorkItem {
423 op_handle: 1,
424 input_handle: 0,
425 output_handle: 5,
426 param: 7,
427 }];
428 let b = vec![MegakernelWorkItem {
429 op_handle: 1,
430 input_handle: 0,
431 output_handle: 5,
432 param: 99,
433 }];
434
435 let arms: [&[MegakernelWorkItem]; 2] = [&a, &b];
436 let report = detect_cross_arm_redundancy(&arms);
437 assert_eq!(report.total_redundant_ops, 1);
438 }
439
440 #[test]
441 fn different_inputs_are_not_redundant() {
442 let a = vec![item(1, 0, 5)];
443 let b = vec![item(1, 1, 5)]; let arms: [&[MegakernelWorkItem]; 2] = [&a, &b];
445 assert!(detect_cross_arm_redundancy(&arms).is_empty());
446 }
447
448 #[test]
449 fn different_outputs_are_not_redundant() {
450 let a = vec![item(1, 0, 5)];
451 let b = vec![item(1, 0, 6)]; let arms: [&[MegakernelWorkItem]; 2] = [&a, &b];
453 assert!(detect_cross_arm_redundancy(&arms).is_empty());
454 }
455
456 #[test]
457 fn op_index_refers_to_late_arm_position() {
458 let a = vec![item(1, 0, 5)];
461 let b = vec![item(99, 0, 0), item(1, 0, 5), item(42, 0, 0)];
462 let arms: [&[MegakernelWorkItem]; 2] = [&a, &b];
463 let report = detect_cross_arm_redundancy(&arms);
464 assert_eq!(report.redundant_pairs, vec![(0, 1, 1)]);
465 }
466
467 #[test]
468 fn prune_redundant_work_items_drops_later_duplicates() {
469 let items = vec![
470 item(1, 0, 5),
471 item(2, 5, 6),
472 item(1, 0, 5),
473 item(3, 6, 7),
474 item(2, 5, 6),
475 ];
476 let mut out = Vec::new();
477
478 let report = prune_redundant_work_items_into(&items, &mut out);
479
480 assert_eq!(out, vec![item(1, 0, 5), item(2, 5, 6), item(3, 6, 7)]);
481 assert_eq!(report.total_redundant_ops, 2);
482 assert_eq!(report.redundant_pairs, vec![(0, 2, 0), (1, 4, 0)]);
483 }
484
485 #[test]
486 fn prune_redundant_work_items_reuses_hash_scratch() {
487 let items = vec![item(1, 0, 5), item(2, 5, 6), item(1, 0, 5), item(3, 6, 7)];
488 let mut out = Vec::new();
489 let mut scratch = RedundantWorkItemPruneScratch::default();
490
491 let first = prune_redundant_work_items_with_scratch_into(&items, &mut out, &mut scratch);
492 let retained_capacity = scratch.first_seen.capacity();
493 out.clear();
494 let second = prune_redundant_work_items_with_scratch_into(&items, &mut out, &mut scratch);
495
496 assert_eq!(first, second);
497 assert!(
498 scratch.first_seen.capacity() >= retained_capacity,
499 "hot megakernel dedupe must retain hash capacity across repeated dispatches"
500 );
501 assert_eq!(out, vec![item(1, 0, 5), item(2, 5, 6), item(3, 6, 7)]);
502 }
503
504 #[test]
505 fn prune_redundant_work_items_handles_empty_input() {
506 let mut out = vec![item(99, 99, 99)];
507
508 let report = prune_redundant_work_items_into(&[], &mut out);
509
510 assert!(report.is_empty());
511 assert!(out.is_empty());
512 }
513
514 #[test]
515 fn prune_redundant_work_items_all_duplicates_keep_one() {
516 let items = vec![item(1, 0, 5), item(1, 0, 5), item(1, 0, 5)];
517 let mut out = Vec::new();
518
519 let report = prune_redundant_work_items_into(&items, &mut out);
520
521 assert_eq!(out, vec![item(1, 0, 5)]);
522 assert_eq!(report.total_redundant_ops, 2);
523 assert_eq!(report.redundant_pairs, vec![(0, 1, 0), (0, 2, 0)]);
524 }
525
526 #[test]
527 fn prune_redundant_work_items_preserves_order_after_first_duplicate() {
528 let items = vec![
529 item(1, 0, 5),
530 item(2, 5, 6),
531 item(1, 0, 5),
532 item(3, 6, 7),
533 item(4, 7, 8),
534 ];
535 let mut out = Vec::new();
536
537 let report = prune_redundant_work_items_into(&items, &mut out);
538
539 assert_eq!(
540 out,
541 vec![item(1, 0, 5), item(2, 5, 6), item(3, 6, 7), item(4, 7, 8)]
542 );
543 assert_eq!(report.redundant_pairs, vec![(0, 2, 0)]);
544 }
545
546 #[test]
547 fn prune_redundant_work_items_leaves_output_empty_when_no_copy_needed() {
548 let items = vec![item(1, 0, 5)];
549 let mut out = vec![item(99, 99, 99)];
550
551 let report = prune_redundant_work_items_into(&items, &mut out);
552
553 assert!(report.is_empty());
554 assert!(out.is_empty());
555 }
556
557 #[test]
558 fn prune_redundant_work_items_keeps_distinct_params() {
559 let mut a = item(1, 0, 5);
560 a.param = 7;
561 let mut b = item(1, 0, 5);
562 b.param = 99;
563 let items = vec![a, b];
564 let mut out = Vec::new();
565
566 let report = prune_redundant_work_items_into(&items, &mut out);
567
568 assert!(report.is_empty());
569 assert!(out.is_empty());
570 }
571
572 #[test]
573 fn output_handles_dense_unique_accepts_single_owner_outputs() {
574 let items = vec![item(1, 0, 5), item(1, 0, 6), item(1, 0, 7)];
575
576 assert!(output_handles_are_dense_unique(&items));
577 }
578
579 #[test]
580 fn output_handles_dense_unique_rejects_repeated_output() {
581 let items = vec![item(1, 0, 5), item(2, 0, 5)];
582
583 assert!(!output_handles_are_dense_unique(&items));
584 }
585
586 #[test]
587 fn prune_redundant_work_items_still_catches_duplicate_with_repeated_output() {
588 let items = vec![item(1, 0, 5), item(2, 0, 6), item(1, 0, 5)];
589 let mut out = Vec::new();
590
591 let report = prune_redundant_work_items_into(&items, &mut out);
592
593 assert_eq!(report.total_redundant_ops, 1);
594 assert_eq!(out, vec![item(1, 0, 5), item(2, 0, 6)]);
595 }
596}