1use std::collections::HashMap;
40use std::fmt;
41use std::sync::{Arc, RwLock};
42use torsh_core::sync::RwLockExt;
43
44use torsh_core::{
45 error::{Result, TorshError},
46 shape::Shape,
47};
48
49#[derive(Debug, Clone, PartialEq, Eq)]
51pub enum ShapeOperation {
52 ElementWise,
54 MatMul,
56 Conv,
58 Pool,
60 Reshape,
62 Transpose,
64 Concatenate,
66 Stack,
68 Broadcast,
70 Reduce,
72 Custom(String),
74}
75
76impl fmt::Display for ShapeOperation {
77 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
78 match self {
79 ShapeOperation::ElementWise => write!(f, "ElementWise"),
80 ShapeOperation::MatMul => write!(f, "MatMul"),
81 ShapeOperation::Conv => write!(f, "Convolution"),
82 ShapeOperation::Pool => write!(f, "Pooling"),
83 ShapeOperation::Reshape => write!(f, "Reshape"),
84 ShapeOperation::Transpose => write!(f, "Transpose"),
85 ShapeOperation::Concatenate => write!(f, "Concatenate"),
86 ShapeOperation::Stack => write!(f, "Stack"),
87 ShapeOperation::Broadcast => write!(f, "Broadcast"),
88 ShapeOperation::Reduce => write!(f, "Reduce"),
89 ShapeOperation::Custom(name) => write!(f, "Custom({})", name),
90 }
91 }
92}
93
94#[derive(Debug, Clone)]
96pub struct ShapeTraceStep {
97 pub step: usize,
99 pub operation: ShapeOperation,
101 pub input_shapes: Vec<Shape>,
103 pub output_shape: Option<Shape>,
105 pub explanation: String,
107 pub success: bool,
109 pub error: Option<String>,
111}
112
113impl fmt::Display for ShapeTraceStep {
114 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
115 writeln!(f, "Step {}: {}", self.step, self.operation)?;
116 writeln!(f, " Input shapes:")?;
117 for (i, shape) in self.input_shapes.iter().enumerate() {
118 writeln!(f, " [{}] {:?}", i, shape.dims())?;
119 }
120 if let Some(output) = &self.output_shape {
121 writeln!(f, " Output shape: {:?}", output.dims())?;
122 }
123 writeln!(f, " Explanation: {}", self.explanation)?;
124 if let Some(error) = &self.error {
125 writeln!(f, " Error: {}", error)?;
126 }
127 writeln!(f, " Status: {}", if self.success { "✓" } else { "✗" })
128 }
129}
130
131#[derive(Debug, Clone)]
133pub struct DebugConfig {
134 pub tracing_enabled: bool,
136 pub auto_validate: bool,
138 pub max_trace_steps: usize,
140 pub include_visuals: bool,
142}
143
144impl Default for DebugConfig {
145 fn default() -> Self {
146 Self {
147 tracing_enabled: true,
148 auto_validate: true,
149 max_trace_steps: 1000,
150 include_visuals: false,
151 }
152 }
153}
154
155pub struct ShapeInferenceDebugger {
157 config: Arc<RwLock<DebugConfig>>,
159 trace: Arc<RwLock<Vec<ShapeTraceStep>>>,
161 step_counter: Arc<RwLock<usize>>,
163 named_shapes: Arc<RwLock<HashMap<String, Shape>>>,
165}
166
167impl ShapeInferenceDebugger {
168 pub fn new() -> Self {
170 Self::with_config(DebugConfig::default())
171 }
172
173 pub fn with_config(config: DebugConfig) -> Self {
175 Self {
176 config: Arc::new(RwLock::new(config)),
177 trace: Arc::new(RwLock::new(Vec::new())),
178 step_counter: Arc::new(RwLock::new(0)),
179 named_shapes: Arc::new(RwLock::new(HashMap::new())),
180 }
181 }
182
183 pub fn enable_tracing(&mut self, enabled: bool) {
185 self.config.write_or_recover().tracing_enabled = enabled;
186 }
187
188 pub fn register_shape(&mut self, name: impl Into<String>, shape: Shape) {
190 self.named_shapes
191 .write_or_recover()
192 .insert(name.into(), shape);
193 }
194
195 fn add_trace_step(&self, step: ShapeTraceStep) {
197 let config = self.config.read_or_recover();
198 if !config.tracing_enabled {
199 return;
200 }
201 drop(config);
202
203 let mut trace = self.trace.write_or_recover();
204 trace.push(step);
205
206 let config = self.config.read_or_recover();
208 if trace.len() > config.max_trace_steps {
209 trace.remove(0);
210 }
211 }
212
213 fn next_step(&self) -> usize {
215 let mut counter = self.step_counter.write_or_recover();
216 let step = *counter;
217 *counter += 1;
218 step
219 }
220
221 pub fn infer_elementwise_shape(&self, shapes: &[Shape]) -> Result<Shape> {
223 let step = self.next_step();
224
225 if shapes.is_empty() {
226 let error = "Element-wise operation requires at least one input".to_string();
227 self.add_trace_step(ShapeTraceStep {
228 step,
229 operation: ShapeOperation::ElementWise,
230 input_shapes: shapes.to_vec(),
231 output_shape: None,
232 explanation: "Checking input shapes".to_string(),
233 success: false,
234 error: Some(error.clone()),
235 });
236 return Err(TorshError::InvalidShape(error));
237 }
238
239 if shapes.len() == 1 {
240 let output = shapes[0].clone();
241 self.add_trace_step(ShapeTraceStep {
242 step,
243 operation: ShapeOperation::ElementWise,
244 input_shapes: shapes.to_vec(),
245 output_shape: Some(output.clone()),
246 explanation: "Single input - output shape matches input".to_string(),
247 success: true,
248 error: None,
249 });
250 return Ok(output);
251 }
252
253 let first_shape = &shapes[0];
255 let all_identical = shapes.iter().all(|s| s.dims() == first_shape.dims());
256
257 if all_identical {
258 let output = first_shape.clone();
259 self.add_trace_step(ShapeTraceStep {
260 step,
261 operation: ShapeOperation::ElementWise,
262 input_shapes: shapes.to_vec(),
263 output_shape: Some(output.clone()),
264 explanation: "All input shapes are identical".to_string(),
265 success: true,
266 error: None,
267 });
268 return Ok(output);
269 }
270
271 let broadcast_result = self.infer_broadcast_shape(shapes);
273 match broadcast_result {
274 Ok(output) => {
275 self.add_trace_step(ShapeTraceStep {
276 step,
277 operation: ShapeOperation::ElementWise,
278 input_shapes: shapes.to_vec(),
279 output_shape: Some(output.clone()),
280 explanation: "Shapes are compatible through broadcasting".to_string(),
281 success: true,
282 error: None,
283 });
284 Ok(output)
285 }
286 Err(e) => {
287 self.add_trace_step(ShapeTraceStep {
288 step,
289 operation: ShapeOperation::ElementWise,
290 input_shapes: shapes.to_vec(),
291 output_shape: None,
292 explanation: "Shapes are not compatible".to_string(),
293 success: false,
294 error: Some(e.to_string()),
295 });
296 Err(e)
297 }
298 }
299 }
300
301 pub fn infer_matmul_shape(&self, a: &Shape, b: &Shape) -> Result<Shape> {
303 let step = self.next_step();
304 let input_shapes = vec![a.clone(), b.clone()];
305
306 let a_dims = a.dims();
307 let b_dims = b.dims();
308
309 if a_dims.len() < 2 || b_dims.len() < 2 {
311 let error = format!(
312 "Matrix multiplication requires at least 2D tensors, got shapes {:?} and {:?}",
313 a_dims, b_dims
314 );
315 self.add_trace_step(ShapeTraceStep {
316 step,
317 operation: ShapeOperation::MatMul,
318 input_shapes,
319 output_shape: None,
320 explanation: "Checking input dimensionality".to_string(),
321 success: false,
322 error: Some(error.clone()),
323 });
324 return Err(TorshError::InvalidShape(error));
325 }
326
327 let a_rows = a_dims[a_dims.len() - 2];
329 let a_cols = a_dims[a_dims.len() - 1];
330 let b_rows = b_dims[b_dims.len() - 2];
331 let b_cols = b_dims[b_dims.len() - 1];
332
333 if a_cols != b_rows {
335 let error = format!(
336 "Matrix multiplication shape mismatch: ({}, {}) @ ({}, {}) - inner dimensions {} and {} do not match",
337 a_rows, a_cols, b_rows, b_cols, a_cols, b_rows
338 );
339 self.add_trace_step(ShapeTraceStep {
340 step,
341 operation: ShapeOperation::MatMul,
342 input_shapes,
343 output_shape: None,
344 explanation: format!(
345 "Checking inner dimension compatibility: {} vs {}",
346 a_cols, b_rows
347 ),
348 success: false,
349 error: Some(error.clone()),
350 });
351 return Err(TorshError::InvalidShape(error));
352 }
353
354 let mut output_dims = Vec::new();
356
357 if a_dims.len() > 2 || b_dims.len() > 2 {
359 let max_batch_dims = std::cmp::max(a_dims.len() - 2, b_dims.len() - 2);
360 for i in 0..max_batch_dims {
361 let a_idx = a_dims.len().saturating_sub(3 + i);
362 let b_idx = b_dims.len().saturating_sub(3 + i);
363
364 let a_dim = if a_idx < a_dims.len() - 2 {
365 a_dims[a_idx]
366 } else {
367 1
368 };
369
370 let b_dim = if b_idx < b_dims.len() - 2 {
371 b_dims[b_idx]
372 } else {
373 1
374 };
375
376 if a_dim != b_dim && a_dim != 1 && b_dim != 1 {
377 let error = format!(
378 "Batch dimension mismatch at position {}: {} vs {}",
379 i, a_dim, b_dim
380 );
381 self.add_trace_step(ShapeTraceStep {
382 step,
383 operation: ShapeOperation::MatMul,
384 input_shapes,
385 output_shape: None,
386 explanation: format!("Checking batch dimension {}", i),
387 success: false,
388 error: Some(error.clone()),
389 });
390 return Err(TorshError::InvalidShape(error));
391 }
392
393 output_dims.push(std::cmp::max(a_dim, b_dim));
394 }
395 }
396
397 output_dims.push(a_rows);
399 output_dims.push(b_cols);
400
401 let output = Shape::new(output_dims);
402
403 self.add_trace_step(ShapeTraceStep {
404 step,
405 operation: ShapeOperation::MatMul,
406 input_shapes,
407 output_shape: Some(output.clone()),
408 explanation: format!(
409 "Matrix multiplication: ({}, {}) @ ({}, {}) = ({}, {})",
410 a_rows, a_cols, b_rows, b_cols, a_rows, b_cols
411 ),
412 success: true,
413 error: None,
414 });
415
416 Ok(output)
417 }
418
419 pub fn infer_broadcast_shape(&self, shapes: &[Shape]) -> Result<Shape> {
421 let step = self.next_step();
422
423 if shapes.is_empty() {
424 let error = "Broadcasting requires at least one shape".to_string();
425 self.add_trace_step(ShapeTraceStep {
426 step,
427 operation: ShapeOperation::Broadcast,
428 input_shapes: shapes.to_vec(),
429 output_shape: None,
430 explanation: "Checking number of inputs".to_string(),
431 success: false,
432 error: Some(error.clone()),
433 });
434 return Err(TorshError::InvalidShape(error));
435 }
436
437 if shapes.len() == 1 {
438 let output = shapes[0].clone();
439 self.add_trace_step(ShapeTraceStep {
440 step,
441 operation: ShapeOperation::Broadcast,
442 input_shapes: shapes.to_vec(),
443 output_shape: Some(output.clone()),
444 explanation: "Single shape - no broadcasting needed".to_string(),
445 success: true,
446 error: None,
447 });
448 return Ok(output);
449 }
450
451 let max_ndim = shapes
453 .iter()
454 .map(|s| s.dims().len())
455 .max()
456 .expect("reduction should succeed");
457
458 let mut output_dims = Vec::with_capacity(max_ndim);
460 let mut explanations = Vec::new();
461
462 for dim_idx in (0..max_ndim).rev() {
463 let mut dim_size = 1;
464 let mut conflict = false;
465 let mut dim_sources = Vec::new();
466
467 for (shape_idx, shape) in shapes.iter().enumerate() {
468 let dims = shape.dims();
469 let pos_from_right = max_ndim - 1 - dim_idx;
471
472 let current_dim = if pos_from_right < dims.len() {
473 dims[dims.len() - 1 - pos_from_right]
474 } else {
475 1
476 };
477
478 if current_dim != 1 {
479 if dim_size == 1 {
480 dim_size = current_dim;
481 dim_sources.push((shape_idx, current_dim));
482 } else if current_dim != dim_size {
483 conflict = true;
484 dim_sources.push((shape_idx, current_dim));
485 }
486 }
487 }
488
489 if conflict {
490 let error = format!(
491 "Broadcasting conflict at dimension {}: incompatible sizes {:?}",
492 dim_idx,
493 dim_sources.iter().map(|(_, d)| d).collect::<Vec<_>>()
494 );
495 self.add_trace_step(ShapeTraceStep {
496 step,
497 operation: ShapeOperation::Broadcast,
498 input_shapes: shapes.to_vec(),
499 output_shape: None,
500 explanation: format!("Checking dimension {}", dim_idx),
501 success: false,
502 error: Some(error.clone()),
503 });
504 return Err(TorshError::InvalidShape(error));
505 }
506
507 output_dims.insert(0, dim_size); if dim_sources.len() > 1 {
509 explanations.insert(
510 0,
511 format!(
512 "Dim {}: broadcast {} sources to size {}",
513 dim_idx,
514 dim_sources.len(),
515 dim_size
516 ),
517 );
518 } else if !dim_sources.is_empty() {
519 explanations.insert(
520 0,
521 format!("Dim {}: size {} (no broadcast)", dim_idx, dim_size),
522 );
523 } else {
524 explanations.insert(0, format!("Dim {}: size 1 (all implicit)", dim_idx));
525 }
526 }
527
528 let output = Shape::new(output_dims);
529 let explanation = if explanations.is_empty() {
530 "All dimensions broadcast successfully".to_string()
531 } else {
532 explanations.join("; ")
533 };
534
535 self.add_trace_step(ShapeTraceStep {
536 step,
537 operation: ShapeOperation::Broadcast,
538 input_shapes: shapes.to_vec(),
539 output_shape: Some(output.clone()),
540 explanation,
541 success: true,
542 error: None,
543 });
544
545 Ok(output)
546 }
547
548 pub fn infer_concat_shape(&self, shapes: &[Shape], dim: i32) -> Result<Shape> {
550 let step = self.next_step();
551
552 if shapes.is_empty() {
553 let error = "Concatenation requires at least one tensor".to_string();
554 self.add_trace_step(ShapeTraceStep {
555 step,
556 operation: ShapeOperation::Concatenate,
557 input_shapes: shapes.to_vec(),
558 output_shape: None,
559 explanation: format!("Checking inputs for concatenation along dim {}", dim),
560 success: false,
561 error: Some(error.clone()),
562 });
563 return Err(TorshError::InvalidShape(error));
564 }
565
566 let first_dims = shapes[0].dims();
567 let ndim = first_dims.len() as i32;
568
569 let concat_dim = if dim < 0 { ndim + dim } else { dim };
571
572 if concat_dim < 0 || concat_dim >= ndim {
573 let error = format!(
574 "Concatenation dimension {} is out of range for {}-dimensional tensors",
575 dim, ndim
576 );
577 self.add_trace_step(ShapeTraceStep {
578 step,
579 operation: ShapeOperation::Concatenate,
580 input_shapes: shapes.to_vec(),
581 output_shape: None,
582 explanation: format!("Validating dimension {}", dim),
583 success: false,
584 error: Some(error.clone()),
585 });
586 return Err(TorshError::InvalidShape(error));
587 }
588
589 let concat_dim = concat_dim as usize;
590 let mut total_concat_size = 0;
591
592 for (i, shape) in shapes.iter().enumerate() {
594 let dims = shape.dims();
595
596 if dims.len() != first_dims.len() {
597 let error = format!(
598 "All tensors must have same number of dimensions for concatenation. \
599 Tensor 0 has {} dims, tensor {} has {} dims",
600 first_dims.len(),
601 i,
602 dims.len()
603 );
604 self.add_trace_step(ShapeTraceStep {
605 step,
606 operation: ShapeOperation::Concatenate,
607 input_shapes: shapes.to_vec(),
608 output_shape: None,
609 explanation: format!("Checking shape compatibility for tensor {}", i),
610 success: false,
611 error: Some(error.clone()),
612 });
613 return Err(TorshError::InvalidShape(error));
614 }
615
616 for (dim_idx, (&d1, &d2)) in first_dims.iter().zip(dims.iter()).enumerate() {
617 if dim_idx != concat_dim && d1 != d2 {
618 let error = format!(
619 "All tensors must have same size in non-concat dimensions. \
620 Dimension {} differs: tensor 0 has size {}, tensor {} has size {}",
621 dim_idx, d1, i, d2
622 );
623 self.add_trace_step(ShapeTraceStep {
624 step,
625 operation: ShapeOperation::Concatenate,
626 input_shapes: shapes.to_vec(),
627 output_shape: None,
628 explanation: format!("Checking dimension {} for tensor {}", dim_idx, i),
629 success: false,
630 error: Some(error.clone()),
631 });
632 return Err(TorshError::InvalidShape(error));
633 }
634 }
635
636 total_concat_size += dims[concat_dim];
637 }
638
639 let mut output_dims = first_dims.to_vec();
641 output_dims[concat_dim] = total_concat_size;
642 let output = Shape::new(output_dims);
643
644 self.add_trace_step(ShapeTraceStep {
645 step,
646 operation: ShapeOperation::Concatenate,
647 input_shapes: shapes.to_vec(),
648 output_shape: Some(output.clone()),
649 explanation: format!(
650 "Concatenating {} tensors along dimension {}: total size {}",
651 shapes.len(),
652 concat_dim,
653 total_concat_size
654 ),
655 success: true,
656 error: None,
657 });
658
659 Ok(output)
660 }
661
662 pub fn get_trace(&self) -> String {
664 let trace = self.trace.read_or_recover();
665 let mut output = String::new();
666 output.push_str("=== Shape Inference Trace ===\n\n");
667
668 for step in trace.iter() {
669 output.push_str(&format!("{}\n", step));
670 }
671
672 output.push_str(&format!("Total steps: {}\n", trace.len()));
673 let successful = trace.iter().filter(|s| s.success).count();
674 output.push_str(&format!(
675 "Successful: {} ({:.1}%)\n",
676 successful,
677 (successful as f64 / trace.len() as f64) * 100.0
678 ));
679
680 output
681 }
682
683 pub fn clear_trace(&mut self) {
685 self.trace.write_or_recover().clear();
686 *self.step_counter.write_or_recover() = 0;
687 }
688
689 pub fn get_statistics(&self) -> TraceStatistics {
691 let trace = self.trace.read_or_recover();
692 let total_steps = trace.len();
693 let successful_steps = trace.iter().filter(|s| s.success).count();
694 let failed_steps = total_steps - successful_steps;
695
696 let mut operation_counts: HashMap<String, usize> = HashMap::new();
697 for step in trace.iter() {
698 *operation_counts
699 .entry(step.operation.to_string())
700 .or_insert(0) += 1;
701 }
702
703 TraceStatistics {
704 total_steps,
705 successful_steps,
706 failed_steps,
707 operation_counts,
708 }
709 }
710}
711
712impl Default for ShapeInferenceDebugger {
713 fn default() -> Self {
714 Self::new()
715 }
716}
717
718#[derive(Debug, Clone)]
720pub struct TraceStatistics {
721 pub total_steps: usize,
723 pub successful_steps: usize,
725 pub failed_steps: usize,
727 pub operation_counts: HashMap<String, usize>,
729}
730
731impl fmt::Display for TraceStatistics {
732 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
733 writeln!(f, "Shape Inference Statistics:")?;
734 writeln!(f, " Total steps: {}", self.total_steps)?;
735 writeln!(f, " Successful: {}", self.successful_steps)?;
736 writeln!(f, " Failed: {}", self.failed_steps)?;
737 if self.total_steps > 0 {
738 writeln!(
739 f,
740 " Success rate: {:.1}%",
741 (self.successful_steps as f64 / self.total_steps as f64) * 100.0
742 )?;
743 }
744 writeln!(f, " Operations:")?;
745 for (op, count) in &self.operation_counts {
746 writeln!(f, " {}: {}", op, count)?;
747 }
748 Ok(())
749 }
750}
751
752#[cfg(test)]
753mod tests {
754 use super::*;
755
756 #[test]
757 fn test_elementwise_identical_shapes() {
758 let debugger = ShapeInferenceDebugger::new();
759 let shapes = vec![
760 Shape::new(vec![2, 3]),
761 Shape::new(vec![2, 3]),
762 Shape::new(vec![2, 3]),
763 ];
764
765 let result = debugger
766 .infer_elementwise_shape(&shapes)
767 .expect("elementwise shape inference should succeed");
768 assert_eq!(result.dims(), &[2, 3]);
769 }
770
771 #[test]
772 fn test_elementwise_broadcasting() {
773 let debugger = ShapeInferenceDebugger::new();
774 let shapes = vec![Shape::new(vec![2, 3]), Shape::new(vec![1, 3])];
775
776 let result = debugger
777 .infer_elementwise_shape(&shapes)
778 .expect("elementwise shape inference should succeed");
779 assert_eq!(result.dims(), &[2, 3]);
780 }
781
782 #[test]
783 fn test_matmul_simple() {
784 let debugger = ShapeInferenceDebugger::new();
785 let a = Shape::new(vec![2, 3]);
786 let b = Shape::new(vec![3, 4]);
787
788 let result = debugger
789 .infer_matmul_shape(&a, &b)
790 .expect("matmul shape inference should succeed");
791 assert_eq!(result.dims(), &[2, 4]);
792 }
793
794 #[test]
795 fn test_matmul_incompatible() {
796 let debugger = ShapeInferenceDebugger::new();
797 let a = Shape::new(vec![2, 3]);
798 let b = Shape::new(vec![4, 5]);
799
800 let result = debugger.infer_matmul_shape(&a, &b);
801 assert!(result.is_err());
802 }
803
804 #[test]
805 fn test_broadcast_compatible() {
806 let debugger = ShapeInferenceDebugger::new();
807 let shapes = vec![
808 Shape::new(vec![2, 3, 4]),
809 Shape::new(vec![3, 4]),
810 Shape::new(vec![4]),
811 ];
812
813 let result = debugger
814 .infer_broadcast_shape(&shapes)
815 .expect("broadcast shape inference should succeed");
816 assert_eq!(result.dims(), &[2, 3, 4]);
817 }
818
819 #[test]
820 fn test_broadcast_incompatible() {
821 let debugger = ShapeInferenceDebugger::new();
822 let shapes = vec![Shape::new(vec![2, 3]), Shape::new(vec![2, 4])];
823
824 let result = debugger.infer_broadcast_shape(&shapes);
825 assert!(result.is_err());
826 }
827
828 #[test]
829 fn test_concat_valid() {
830 let debugger = ShapeInferenceDebugger::new();
831 let shapes = vec![
832 Shape::new(vec![2, 3]),
833 Shape::new(vec![2, 5]),
834 Shape::new(vec![2, 2]),
835 ];
836
837 let result = debugger
838 .infer_concat_shape(&shapes, 1)
839 .expect("concat shape inference should succeed");
840 assert_eq!(result.dims(), &[2, 10]); }
842
843 #[test]
844 fn test_concat_incompatible() {
845 let debugger = ShapeInferenceDebugger::new();
846 let shapes = vec![Shape::new(vec![2, 3]), Shape::new(vec![3, 3])];
847
848 let result = debugger.infer_concat_shape(&shapes, 1);
849 assert!(result.is_err());
850 }
851
852 #[test]
853 fn test_trace_collection() {
854 let debugger = ShapeInferenceDebugger::new();
855
856 let a = Shape::new(vec![2, 3]);
857 let b = Shape::new(vec![3, 4]);
858 let _ = debugger
859 .infer_matmul_shape(&a, &b)
860 .expect("matmul shape inference should succeed");
861
862 let trace = debugger.get_trace();
863 assert!(trace.contains("MatMul"));
864 assert!(trace.contains("Step 0"));
865 }
866
867 #[test]
868 fn test_statistics() {
869 let debugger = ShapeInferenceDebugger::new();
870
871 let _ = debugger.infer_matmul_shape(&Shape::new(vec![2, 3]), &Shape::new(vec![3, 4]));
873 let _ = debugger.infer_broadcast_shape(&[Shape::new(vec![2, 3]), Shape::new(vec![3])]);
874
875 let _ = debugger.infer_matmul_shape(&Shape::new(vec![2, 3]), &Shape::new(vec![4, 5]));
877
878 let stats = debugger.get_statistics();
879 assert_eq!(stats.total_steps, 3);
880 assert_eq!(stats.successful_steps, 2);
881 assert_eq!(stats.failed_steps, 1);
882 }
883}