Skip to main content

torsh_tensor/
shape_inference_debugger.rs

1//! Shape Inference Debugging with Detailed Traces
2//!
3//! This module provides comprehensive debugging capabilities for shape inference operations,
4//! helping developers understand how tensor shapes are computed and why shape mismatches occur.
5//!
6//! # Features
7//!
8//! - **Detailed trace logging**: Record every step of shape inference with explanations
9//! - **Shape compatibility checking**: Validate shape compatibility with detailed error messages
10//! - **Broadcasting visualization**: Visualize how broadcasting affects shapes
11//! - **Shape transformation tracking**: Track how shapes change through operations
12//! - **Interactive debugging**: Step-by-step shape inference with intermediate results
13//! - **Error diagnosis**: Detailed error messages explaining shape mismatches
14//!
15//! # Example
16//!
17//! ```rust
18//! use torsh_tensor::shape_inference_debugger::*;
19//! use torsh_core::shape::Shape;
20//!
21//! # fn main() -> Result<(), Box<dyn std::error::Error>> {
22//! let mut debugger = ShapeInferenceDebugger::new();
23//!
24//! // Enable tracing
25//! debugger.enable_tracing(true);
26//!
27//! // Infer shape for matmul operation
28//! let a_shape = Shape::new(vec![2, 3]);
29//! let b_shape = Shape::new(vec![3, 4]);
30//! let result_shape = debugger.infer_matmul_shape(&a_shape, &b_shape)?;
31//!
32//! // Get detailed trace
33//! let trace = debugger.get_trace();
34//! println!("Shape inference trace:\n{}", trace);
35//! # Ok(())
36//! # }
37//! ```
38
39use 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/// Type of shape operation
50#[derive(Debug, Clone, PartialEq, Eq)]
51pub enum ShapeOperation {
52    /// Element-wise operation
53    ElementWise,
54    /// Matrix multiplication
55    MatMul,
56    /// Convolution operation
57    Conv,
58    /// Pooling operation
59    Pool,
60    /// Reshape operation
61    Reshape,
62    /// Transpose operation
63    Transpose,
64    /// Concatenation
65    Concatenate,
66    /// Stack operation
67    Stack,
68    /// Broadcast operation
69    Broadcast,
70    /// Reduction operation
71    Reduce,
72    /// Custom operation
73    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/// A single step in shape inference trace
95#[derive(Debug, Clone)]
96pub struct ShapeTraceStep {
97    /// Step number
98    pub step: usize,
99    /// Operation being performed
100    pub operation: ShapeOperation,
101    /// Input shapes
102    pub input_shapes: Vec<Shape>,
103    /// Output shape
104    pub output_shape: Option<Shape>,
105    /// Explanation of the inference
106    pub explanation: String,
107    /// Whether this step succeeded
108    pub success: bool,
109    /// Error message if failed
110    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/// Configuration for shape inference debugging
132#[derive(Debug, Clone)]
133pub struct DebugConfig {
134    /// Whether to enable detailed tracing
135    pub tracing_enabled: bool,
136    /// Whether to validate shapes automatically
137    pub auto_validate: bool,
138    /// Maximum number of trace steps to keep
139    pub max_trace_steps: usize,
140    /// Whether to include visual diagrams
141    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
155/// Main shape inference debugger
156pub struct ShapeInferenceDebugger {
157    /// Configuration
158    config: Arc<RwLock<DebugConfig>>,
159    /// Trace of shape inference steps
160    trace: Arc<RwLock<Vec<ShapeTraceStep>>>,
161    /// Current step counter
162    step_counter: Arc<RwLock<usize>>,
163    /// Named shapes for reference
164    named_shapes: Arc<RwLock<HashMap<String, Shape>>>,
165}
166
167impl ShapeInferenceDebugger {
168    /// Create a new shape inference debugger
169    pub fn new() -> Self {
170        Self::with_config(DebugConfig::default())
171    }
172
173    /// Create a debugger with custom configuration
174    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    /// Enable or disable tracing
184    pub fn enable_tracing(&mut self, enabled: bool) {
185        self.config.write_or_recover().tracing_enabled = enabled;
186    }
187
188    /// Register a named shape for reference
189    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    /// Add a trace step
196    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        // Trim if needed
207        let config = self.config.read_or_recover();
208        if trace.len() > config.max_trace_steps {
209            trace.remove(0);
210        }
211    }
212
213    /// Get the next step number
214    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    /// Infer shape for element-wise operation
222    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        // Check if all shapes are identical
254        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        // Try broadcasting
272        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    /// Infer shape for matrix multiplication
302    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        // Check dimensions
310        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        // Get the matrix dimensions
328        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        // Check compatibility
334        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        // Build output shape
355        let mut output_dims = Vec::new();
356
357        // Handle batch dimensions
358        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        // Add matrix dimensions
398        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    /// Infer shape for broadcasting
420    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        // Find maximum number of dimensions
452        let max_ndim = shapes
453            .iter()
454            .map(|s| s.dims().len())
455            .max()
456            .expect("reduction should succeed");
457
458        // Build output shape by checking each dimension (from right to left for broadcasting)
459        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                // Access dimensions from the right
470                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); // Insert at front since we're going right to left
508            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    /// Infer shape for concatenation
549    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        // Normalize dimension
570        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        // Check all shapes are compatible
593        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        // Build output shape
640        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    /// Get the complete trace as a formatted string
663    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    /// Clear the trace
684    pub fn clear_trace(&mut self) {
685        self.trace.write_or_recover().clear();
686        *self.step_counter.write_or_recover() = 0;
687    }
688
689    /// Get trace statistics
690    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/// Statistics about the shape inference trace
719#[derive(Debug, Clone)]
720pub struct TraceStatistics {
721    /// Total number of trace steps
722    pub total_steps: usize,
723    /// Number of successful steps
724    pub successful_steps: usize,
725    /// Number of failed steps
726    pub failed_steps: usize,
727    /// Count of each operation type
728    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]); // 3 + 5 + 2
841    }
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        // Successful operations
872        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        // Failed operation
876        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}