Skip to main content

torsh_autograd/
function_optimization.rs

1//! Function optimization and fusion framework
2//!
3//! This module provides advanced optimization techniques for autograd functions,
4//! including operation fusion, pattern matching, and performance optimization.
5
6use crate::function::FunctionMetadata;
7use std::collections::HashMap;
8use torsh_core::error::{Result, TorshError};
9use tracing::{debug, info};
10
11/// Optimization strategies for autograd functions
12#[derive(Debug, Clone, Copy, PartialEq)]
13pub enum OptimizationStrategy {
14    /// Fuse sequential operations
15    SequentialFusion,
16    /// Fuse element-wise operations
17    ElementWiseFusion,
18    /// Fuse matrix operations
19    MatrixFusion,
20    /// Eliminate common subexpressions
21    CommonSubexpressionElimination,
22    /// Dead code elimination
23    DeadCodeElimination,
24    /// Constant folding
25    ConstantFolding,
26    /// Memory layout optimization
27    MemoryLayoutOptimization,
28    /// SIMD vectorization
29    SIMDVectorization,
30}
31
32/// Configuration for function optimization
33#[derive(Debug, Clone)]
34pub struct OptimizationConfig {
35    /// Enabled optimization strategies
36    pub enabled_strategies: Vec<OptimizationStrategy>,
37    /// Maximum fusion group size
38    pub max_fusion_size: usize,
39    /// Memory threshold for optimizations
40    pub memory_threshold: usize,
41    /// Compute threshold for optimizations
42    pub compute_threshold: f32,
43    /// Enable aggressive optimizations
44    pub aggressive_mode: bool,
45    /// Profile-guided optimization
46    pub enable_pgo: bool,
47}
48
49impl Default for OptimizationConfig {
50    fn default() -> Self {
51        Self {
52            enabled_strategies: vec![
53                OptimizationStrategy::SequentialFusion,
54                OptimizationStrategy::ElementWiseFusion,
55                OptimizationStrategy::CommonSubexpressionElimination,
56                OptimizationStrategy::DeadCodeElimination,
57            ],
58            max_fusion_size: 8,
59            memory_threshold: 64 * 1024 * 1024, // 64MB
60            compute_threshold: 1.5,             // 50% overhead max
61            aggressive_mode: false,
62            enable_pgo: false,
63        }
64    }
65}
66
67/// Function pattern for optimization matching
68#[derive(Debug, Clone)]
69pub struct FunctionPattern {
70    /// Pattern name
71    pub name: String,
72    /// Sequence of operation names that match this pattern
73    pub operations: Vec<String>,
74    /// Required properties for matching
75    pub required_properties: Vec<PatternProperty>,
76    /// Optimization strategy to apply
77    pub optimization: OptimizationStrategy,
78    /// Priority for pattern matching
79    pub priority: i32,
80}
81
82/// Properties required for pattern matching
83#[derive(Debug, Clone, PartialEq)]
84pub enum PatternProperty {
85    /// Operations must be consecutive
86    Consecutive,
87    /// Operations must have compatible shapes
88    CompatibleShapes,
89    /// Operations must be element-wise
90    ElementWise,
91    /// Operations must be commutative
92    Commutative,
93    /// Operations must have no side effects
94    NoSideEffects,
95}
96
97/// Fusable operation group
98#[derive(Debug, Clone)]
99pub struct FusionGroup {
100    /// Operations in the fusion group
101    pub operations: Vec<FunctionInfo>,
102    /// Fusion strategy to apply
103    pub strategy: OptimizationStrategy,
104    /// Estimated performance gain
105    pub performance_gain: f32,
106    /// Estimated memory savings
107    pub memory_savings: usize,
108}
109
110/// Information about a function for optimization
111#[derive(Debug, Clone)]
112pub struct FunctionInfo {
113    /// Function identifier
114    pub id: usize,
115    /// Function name
116    pub name: String,
117    /// Function metadata
118    pub metadata: FunctionMetadata,
119    /// Input shapes (if known)
120    pub input_shapes: Vec<Vec<usize>>,
121    /// Output shapes (if known)
122    pub output_shapes: Vec<Vec<usize>>,
123    /// Dependencies
124    pub dependencies: Vec<usize>,
125    /// Performance profile data
126    pub profile_data: Option<ProfileData>,
127}
128
129/// Performance profiling data for functions
130#[derive(Debug, Clone)]
131pub struct ProfileData {
132    /// Average execution time (milliseconds)
133    pub avg_execution_time: f32,
134    /// Memory usage (bytes)
135    pub memory_usage: usize,
136    /// Cache hit rate
137    pub cache_hit_rate: f32,
138    /// Number of executions
139    pub execution_count: usize,
140}
141
142/// Function optimizer and fusion engine
143pub struct FunctionOptimizer {
144    config: OptimizationConfig,
145    patterns: Vec<FunctionPattern>,
146    function_registry: HashMap<usize, FunctionInfo>,
147    optimization_history: Vec<OptimizationResult>,
148    #[allow(dead_code)]
149    profile_database: HashMap<String, ProfileData>,
150}
151
152/// Result of an optimization pass
153#[derive(Debug, Clone)]
154pub struct OptimizationResult {
155    /// Strategy that was applied
156    pub strategy: OptimizationStrategy,
157    /// Functions that were optimized
158    pub optimized_functions: Vec<usize>,
159    /// Performance improvement estimate
160    pub performance_improvement: f32,
161    /// Memory savings estimate
162    pub memory_savings: usize,
163    /// Whether optimization was successful
164    pub success: bool,
165    /// Error message if optimization failed
166    pub error_message: Option<String>,
167}
168
169impl FunctionOptimizer {
170    /// Create a new function optimizer
171    pub fn new(config: OptimizationConfig) -> Self {
172        let mut optimizer = Self {
173            config,
174            patterns: Vec::new(),
175            function_registry: HashMap::new(),
176            optimization_history: Vec::new(),
177            profile_database: HashMap::new(),
178        };
179
180        optimizer.initialize_default_patterns();
181        optimizer
182    }
183
184    /// Initialize default optimization patterns
185    fn initialize_default_patterns(&mut self) {
186        // Element-wise fusion patterns
187        self.patterns.push(FunctionPattern {
188            name: "ElementWise_Add_Mul".to_string(),
189            operations: vec!["add".to_string(), "mul".to_string()],
190            required_properties: vec![
191                PatternProperty::Consecutive,
192                PatternProperty::ElementWise,
193                PatternProperty::CompatibleShapes,
194            ],
195            optimization: OptimizationStrategy::ElementWiseFusion,
196            priority: 100,
197        });
198
199        self.patterns.push(FunctionPattern {
200            name: "ElementWise_Chain".to_string(),
201            operations: vec!["relu".to_string(), "mul".to_string(), "add".to_string()],
202            required_properties: vec![PatternProperty::Consecutive, PatternProperty::ElementWise],
203            optimization: OptimizationStrategy::ElementWiseFusion,
204            priority: 90,
205        });
206
207        // Matrix operation fusion patterns
208        self.patterns.push(FunctionPattern {
209            name: "MatMul_Add".to_string(),
210            operations: vec!["matmul".to_string(), "add".to_string()],
211            required_properties: vec![
212                PatternProperty::Consecutive,
213                PatternProperty::CompatibleShapes,
214            ],
215            optimization: OptimizationStrategy::MatrixFusion,
216            priority: 95,
217        });
218
219        // Sequential operation patterns
220        self.patterns.push(FunctionPattern {
221            name: "Sequential_Activation".to_string(),
222            operations: vec!["linear".to_string(), "relu".to_string()],
223            required_properties: vec![PatternProperty::Consecutive, PatternProperty::NoSideEffects],
224            optimization: OptimizationStrategy::SequentialFusion,
225            priority: 85,
226        });
227
228        info!("Initialized {} optimization patterns", self.patterns.len());
229    }
230
231    /// Register a function for optimization tracking
232    pub fn register_function(&mut self, function_info: FunctionInfo) {
233        self.function_registry
234            .insert(function_info.id, function_info);
235    }
236
237    /// Optimize a sequence of functions
238    pub fn optimize_functions(
239        &mut self,
240        function_ids: &[usize],
241    ) -> Result<Vec<OptimizationResult>> {
242        let mut results = Vec::new();
243
244        for &strategy in &self.config.enabled_strategies.clone() {
245            if let Ok(result) = self.apply_optimization_strategy(strategy, function_ids) {
246                results.push(result);
247            }
248        }
249
250        // Store optimization history
251        self.optimization_history.extend(results.iter().cloned());
252
253        info!(
254            "Applied {} optimization strategies to {} functions",
255            results.len(),
256            function_ids.len()
257        );
258
259        Ok(results)
260    }
261
262    /// Apply a specific optimization strategy
263    fn apply_optimization_strategy(
264        &mut self,
265        strategy: OptimizationStrategy,
266        function_ids: &[usize],
267    ) -> Result<OptimizationResult> {
268        match strategy {
269            OptimizationStrategy::SequentialFusion => self.apply_sequential_fusion(function_ids),
270            OptimizationStrategy::ElementWiseFusion => self.apply_element_wise_fusion(function_ids),
271            OptimizationStrategy::MatrixFusion => self.apply_matrix_fusion(function_ids),
272            OptimizationStrategy::CommonSubexpressionElimination => self.apply_cse(function_ids),
273            OptimizationStrategy::DeadCodeElimination => self.apply_dce(function_ids),
274            OptimizationStrategy::ConstantFolding => self.apply_constant_folding(function_ids),
275            OptimizationStrategy::MemoryLayoutOptimization => {
276                self.apply_memory_layout_optimization(function_ids)
277            }
278            OptimizationStrategy::SIMDVectorization => self.apply_simd_vectorization(function_ids),
279        }
280    }
281
282    /// Apply sequential fusion optimization
283    fn apply_sequential_fusion(&mut self, function_ids: &[usize]) -> Result<OptimizationResult> {
284        let fusion_groups = self.find_sequential_fusion_opportunities(function_ids)?;
285
286        let mut optimized_functions = Vec::new();
287        let mut total_performance_gain = 0.0;
288        let mut total_memory_savings = 0;
289
290        for group in fusion_groups {
291            // Apply fusion to the group
292            let _fusion_result = self.fuse_sequential_operations(&group)?;
293
294            optimized_functions.extend(group.operations.iter().map(|op| op.id));
295            total_performance_gain += group.performance_gain;
296            total_memory_savings += group.memory_savings;
297
298            debug!(
299                "Fused {} sequential operations with {:.2}% performance gain",
300                group.operations.len(),
301                group.performance_gain * 100.0
302            );
303        }
304
305        Ok(OptimizationResult {
306            strategy: OptimizationStrategy::SequentialFusion,
307            optimized_functions,
308            performance_improvement: total_performance_gain,
309            memory_savings: total_memory_savings,
310            success: true,
311            error_message: None,
312        })
313    }
314
315    /// Apply element-wise fusion optimization
316    fn apply_element_wise_fusion(&mut self, function_ids: &[usize]) -> Result<OptimizationResult> {
317        let fusion_groups = self.find_element_wise_fusion_opportunities(function_ids)?;
318
319        let mut optimized_functions = Vec::new();
320        let mut total_performance_gain = 0.0;
321        let mut total_memory_savings = 0;
322
323        for group in fusion_groups {
324            optimized_functions.extend(group.operations.iter().map(|op| op.id));
325            total_performance_gain += group.performance_gain;
326            total_memory_savings += group.memory_savings;
327
328            debug!("Fused {} element-wise operations", group.operations.len());
329        }
330
331        Ok(OptimizationResult {
332            strategy: OptimizationStrategy::ElementWiseFusion,
333            optimized_functions,
334            performance_improvement: total_performance_gain,
335            memory_savings: total_memory_savings,
336            success: true,
337            error_message: None,
338        })
339    }
340
341    /// Apply matrix fusion optimization
342    fn apply_matrix_fusion(&mut self, function_ids: &[usize]) -> Result<OptimizationResult> {
343        let mut optimized_functions = Vec::new();
344        let mut performance_gain = 0.0;
345        let mut memory_savings = 0;
346
347        // Look for matrix multiplication followed by addition (GEMM pattern)
348        for window in function_ids.windows(2) {
349            if let (Some(func1), Some(func2)) = (
350                self.function_registry.get(&window[0]),
351                self.function_registry.get(&window[1]),
352            ) {
353                if func1.name == "matmul" && func2.name == "add" {
354                    // This is a fusable GEMM pattern
355                    optimized_functions.extend_from_slice(window);
356                    performance_gain += 0.2; // 20% improvement estimate
357                    memory_savings += 1024 * 1024; // 1MB savings estimate
358
359                    debug!("Fused MatMul+Add pattern");
360                }
361            }
362        }
363
364        let success = !optimized_functions.is_empty();
365        Ok(OptimizationResult {
366            strategy: OptimizationStrategy::MatrixFusion,
367            optimized_functions,
368            performance_improvement: performance_gain,
369            memory_savings,
370            success,
371            error_message: None,
372        })
373    }
374
375    /// Apply common subexpression elimination
376    fn apply_cse(&mut self, function_ids: &[usize]) -> Result<OptimizationResult> {
377        let mut optimized_functions = Vec::new();
378        let mut expression_map: HashMap<String, Vec<usize>> = HashMap::new();
379
380        // Group functions by their "signature" (name + input shapes)
381        for &func_id in function_ids {
382            if let Some(func_info) = self.function_registry.get(&func_id) {
383                let signature = format!("{}_{:?}", func_info.name, func_info.input_shapes);
384                expression_map.entry(signature).or_default().push(func_id);
385            }
386        }
387
388        let mut eliminated_expressions = 0;
389
390        // Find common subexpressions
391        for (signature, func_ids) in expression_map {
392            if func_ids.len() > 1 {
393                // Found common subexpression
394                optimized_functions.extend(&func_ids[1..]); // Keep first, eliminate rest
395                eliminated_expressions += func_ids.len() - 1;
396
397                debug!(
398                    "Eliminated {} instances of common subexpression: {}",
399                    func_ids.len() - 1,
400                    signature
401                );
402            }
403        }
404
405        Ok(OptimizationResult {
406            strategy: OptimizationStrategy::CommonSubexpressionElimination,
407            optimized_functions,
408            performance_improvement: eliminated_expressions as f32 * 0.1, // 10% per elimination
409            memory_savings: eliminated_expressions * 1024,                // 1KB per elimination
410            success: eliminated_expressions > 0,
411            error_message: None,
412        })
413    }
414
415    /// Apply dead code elimination
416    fn apply_dce(&mut self, function_ids: &[usize]) -> Result<OptimizationResult> {
417        let mut dead_functions = Vec::new();
418
419        // Simple DCE: find functions with no dependencies in subsequent operations
420        for &func_id in function_ids {
421            if self.is_function_dead(func_id, function_ids) {
422                dead_functions.push(func_id);
423            }
424        }
425
426        debug!(
427            "Found {} dead functions for elimination",
428            dead_functions.len()
429        );
430
431        Ok(OptimizationResult {
432            strategy: OptimizationStrategy::DeadCodeElimination,
433            optimized_functions: dead_functions.clone(),
434            performance_improvement: dead_functions.len() as f32 * 0.05, // 5% per dead function
435            memory_savings: dead_functions.len() * 512,                  // 512B per dead function
436            success: !dead_functions.is_empty(),
437            error_message: None,
438        })
439    }
440
441    /// Apply constant folding optimization
442    fn apply_constant_folding(&mut self, function_ids: &[usize]) -> Result<OptimizationResult> {
443        let mut folded_functions = Vec::new();
444
445        // Look for operations with constant inputs
446        for &func_id in function_ids {
447            if let Some(func_info) = self.function_registry.get(&func_id) {
448                if self.can_constant_fold(&func_info) {
449                    folded_functions.push(func_id);
450                }
451            }
452        }
453
454        debug!(
455            "Found {} functions for constant folding",
456            folded_functions.len()
457        );
458
459        Ok(OptimizationResult {
460            strategy: OptimizationStrategy::ConstantFolding,
461            optimized_functions: folded_functions.clone(),
462            performance_improvement: folded_functions.len() as f32 * 0.3, // 30% per folded function
463            memory_savings: folded_functions.len() * 256, // 256B per folded function
464            success: !folded_functions.is_empty(),
465            error_message: None,
466        })
467    }
468
469    /// Apply memory layout optimization
470    fn apply_memory_layout_optimization(
471        &mut self,
472        function_ids: &[usize],
473    ) -> Result<OptimizationResult> {
474        let mut optimized_functions = Vec::new();
475
476        // Look for opportunities to reorder operations for better memory access patterns
477        for &func_id in function_ids {
478            if let Some(func_info) = self.function_registry.get(&func_id) {
479                if self.can_optimize_memory_layout(&func_info) {
480                    optimized_functions.push(func_id);
481                }
482            }
483        }
484
485        debug!(
486            "Found {} functions for memory layout optimization",
487            optimized_functions.len()
488        );
489
490        Ok(OptimizationResult {
491            strategy: OptimizationStrategy::MemoryLayoutOptimization,
492            optimized_functions: optimized_functions.clone(),
493            performance_improvement: optimized_functions.len() as f32 * 0.15, // 15% per optimization
494            memory_savings: 0, // Memory layout doesn't save memory, but improves access patterns
495            success: !optimized_functions.is_empty(),
496            error_message: None,
497        })
498    }
499
500    /// Apply SIMD vectorization
501    fn apply_simd_vectorization(&mut self, function_ids: &[usize]) -> Result<OptimizationResult> {
502        let mut vectorized_functions = Vec::new();
503
504        // Look for element-wise operations that can be vectorized
505        for &func_id in function_ids {
506            if let Some(func_info) = self.function_registry.get(&func_id) {
507                if self.can_vectorize(&func_info) {
508                    vectorized_functions.push(func_id);
509                }
510            }
511        }
512
513        debug!(
514            "Found {} functions for SIMD vectorization",
515            vectorized_functions.len()
516        );
517
518        Ok(OptimizationResult {
519            strategy: OptimizationStrategy::SIMDVectorization,
520            optimized_functions: vectorized_functions.clone(),
521            performance_improvement: vectorized_functions.len() as f32 * 0.4, // 40% per vectorized function
522            memory_savings: 0, // SIMD doesn't save memory but improves throughput
523            success: !vectorized_functions.is_empty(),
524            error_message: None,
525        })
526    }
527
528    /// Find sequential fusion opportunities
529    fn find_sequential_fusion_opportunities(
530        &self,
531        function_ids: &[usize],
532    ) -> Result<Vec<FusionGroup>> {
533        let mut fusion_groups = Vec::new();
534
535        // Use a sliding window to find consecutive operations that can be fused
536        for window_size in (2..=self.config.max_fusion_size).rev() {
537            for window in function_ids.windows(window_size) {
538                if let Some(group) =
539                    self.analyze_fusion_group(window, OptimizationStrategy::SequentialFusion)?
540                {
541                    fusion_groups.push(group);
542                }
543            }
544        }
545
546        Ok(fusion_groups)
547    }
548
549    /// Find element-wise fusion opportunities
550    fn find_element_wise_fusion_opportunities(
551        &self,
552        function_ids: &[usize],
553    ) -> Result<Vec<FusionGroup>> {
554        let mut fusion_groups = Vec::new();
555
556        // Look for chains of element-wise operations
557        let mut current_group = Vec::new();
558
559        for &func_id in function_ids {
560            if let Some(func_info) = self.function_registry.get(&func_id) {
561                if self.is_element_wise_operation(&func_info) {
562                    current_group.push(func_info.clone());
563                } else {
564                    if current_group.len() >= 2 {
565                        if let Some(group) = self.create_fusion_group(
566                            current_group.clone(),
567                            OptimizationStrategy::ElementWiseFusion,
568                        )? {
569                            fusion_groups.push(group);
570                        }
571                    }
572                    current_group.clear();
573                }
574            }
575        }
576
577        // Handle remaining group
578        if current_group.len() >= 2 {
579            if let Some(group) =
580                self.create_fusion_group(current_group, OptimizationStrategy::ElementWiseFusion)?
581            {
582                fusion_groups.push(group);
583            }
584        }
585
586        Ok(fusion_groups)
587    }
588
589    /// Analyze a potential fusion group
590    fn analyze_fusion_group(
591        &self,
592        function_ids: &[usize],
593        strategy: OptimizationStrategy,
594    ) -> Result<Option<FusionGroup>> {
595        let operations: Result<Vec<_>> = function_ids
596            .iter()
597            .map(|&id| {
598                self.function_registry
599                    .get(&id)
600                    .ok_or_else(|| TorshError::AutogradError(format!("Function {} not found", id)))
601                    .map(|f| f.clone())
602            })
603            .collect();
604
605        let operations = operations?;
606        self.create_fusion_group(operations, strategy)
607    }
608
609    /// Create a fusion group from operations
610    fn create_fusion_group(
611        &self,
612        operations: Vec<FunctionInfo>,
613        strategy: OptimizationStrategy,
614    ) -> Result<Option<FusionGroup>> {
615        if operations.len() < 2 {
616            return Ok(None);
617        }
618
619        // Estimate performance gain and memory savings
620        let performance_gain = self.estimate_fusion_performance_gain(&operations, strategy);
621        let memory_savings = self.estimate_fusion_memory_savings(&operations, strategy);
622
623        // Only create fusion group if it's beneficial
624        if performance_gain > 0.05 || memory_savings > 1024 {
625            // 5% gain or 1KB savings
626            Ok(Some(FusionGroup {
627                operations,
628                strategy,
629                performance_gain,
630                memory_savings,
631            }))
632        } else {
633            Ok(None)
634        }
635    }
636
637    /// Estimate performance gain from fusion
638    fn estimate_fusion_performance_gain(
639        &self,
640        operations: &[FunctionInfo],
641        strategy: OptimizationStrategy,
642    ) -> f32 {
643        let base_gain = match strategy {
644            OptimizationStrategy::SequentialFusion => 0.1,
645            OptimizationStrategy::ElementWiseFusion => 0.2,
646            OptimizationStrategy::MatrixFusion => 0.3,
647            _ => 0.05,
648        };
649
650        base_gain * (operations.len() - 1) as f32
651    }
652
653    /// Estimate memory savings from fusion
654    fn estimate_fusion_memory_savings(
655        &self,
656        operations: &[FunctionInfo],
657        _strategy: OptimizationStrategy,
658    ) -> usize {
659        // Rough estimate: save intermediate results
660        (operations.len() - 1) * 1024 // 1KB per eliminated intermediate
661    }
662
663    /// Check if a function is element-wise
664    fn is_element_wise_operation(&self, func_info: &FunctionInfo) -> bool {
665        matches!(
666            func_info.name.as_str(),
667            "add" | "mul" | "sub" | "div" | "relu" | "sigmoid" | "tanh"
668        )
669    }
670
671    /// Check if a function is dead (unused)
672    fn is_function_dead(&self, func_id: usize, all_functions: &[usize]) -> bool {
673        // Simple check: if no other function depends on this one
674        for &other_id in all_functions {
675            if other_id != func_id {
676                if let Some(other_func) = self.function_registry.get(&other_id) {
677                    if other_func.dependencies.contains(&func_id) {
678                        return false;
679                    }
680                }
681            }
682        }
683        true
684    }
685
686    /// Check if a function can be constant folded
687    fn can_constant_fold(&self, func_info: &FunctionInfo) -> bool {
688        // Simplified check: functions with no dependencies might be constant
689        func_info.dependencies.is_empty()
690            && matches!(func_info.name.as_str(), "constant" | "zeros" | "ones")
691    }
692
693    /// Check if memory layout can be optimized
694    fn can_optimize_memory_layout(&self, func_info: &FunctionInfo) -> bool {
695        // Look for operations that would benefit from memory reordering
696        matches!(func_info.name.as_str(), "transpose" | "reshape" | "permute")
697    }
698
699    /// Check if function can be vectorized
700    fn can_vectorize(&self, func_info: &FunctionInfo) -> bool {
701        // Element-wise operations are good candidates for vectorization
702        self.is_element_wise_operation(func_info)
703            && func_info
704                .input_shapes
705                .iter()
706                .any(|shape| shape.iter().product::<usize>() > 64)
707    }
708
709    /// Actually fuse sequential operations (placeholder)
710    fn fuse_sequential_operations(&mut self, _group: &FusionGroup) -> Result<()> {
711        // In a real implementation, this would create a new fused function
712        // and update the computation graph
713        Ok(())
714    }
715
716    /// Get optimization statistics
717    pub fn get_optimization_stats(&self) -> OptimizationStats {
718        let total_optimizations = self.optimization_history.len();
719        let successful_optimizations = self
720            .optimization_history
721            .iter()
722            .filter(|r| r.success)
723            .count();
724
725        let total_performance_gain: f32 = self
726            .optimization_history
727            .iter()
728            .map(|r| r.performance_improvement)
729            .sum();
730
731        let total_memory_savings: usize = self
732            .optimization_history
733            .iter()
734            .map(|r| r.memory_savings)
735            .sum();
736
737        OptimizationStats {
738            total_optimizations,
739            successful_optimizations,
740            success_rate: if total_optimizations > 0 {
741                successful_optimizations as f32 / total_optimizations as f32
742            } else {
743                0.0
744            },
745            total_performance_gain,
746            total_memory_savings,
747            registered_functions: self.function_registry.len(),
748            active_patterns: self.patterns.len(),
749        }
750    }
751}
752
753/// Statistics about optimization performance
754#[derive(Debug, Clone)]
755pub struct OptimizationStats {
756    pub total_optimizations: usize,
757    pub successful_optimizations: usize,
758    pub success_rate: f32,
759    pub total_performance_gain: f32,
760    pub total_memory_savings: usize,
761    pub registered_functions: usize,
762    pub active_patterns: usize,
763}
764
765#[cfg(test)]
766mod tests {
767    use super::*;
768    use crate::function::{ComputationalComplexity, MemoryComplexity};
769
770    #[test]
771    fn test_function_optimizer_creation() {
772        let config = OptimizationConfig::default();
773        let optimizer = FunctionOptimizer::new(config);
774
775        assert!(!optimizer.patterns.is_empty());
776        assert!(optimizer.function_registry.is_empty());
777    }
778
779    #[test]
780    fn test_pattern_matching() {
781        let config = OptimizationConfig::default();
782        let mut optimizer = FunctionOptimizer::new(config);
783
784        // Register some functions
785        let func1 = FunctionInfo {
786            id: 1,
787            name: "add".to_string(),
788            metadata: FunctionMetadata {
789                name: "add".to_string(),
790                is_differentiable: true,
791                memory_complexity: MemoryComplexity::Linear,
792                computational_complexity: ComputationalComplexity::Linear,
793                is_fusable: true,
794                version: "1.0.0".to_string(),
795                description: "Element-wise addition".to_string(),
796                author: "torsh-autograd".to_string(),
797                created_at: "2024-01-01T00:00:00Z".to_string(),
798                checksum: "".to_string(),
799                dependencies: vec![],
800            },
801            input_shapes: vec![vec![10, 10]],
802            output_shapes: vec![vec![10, 10]],
803            dependencies: vec![],
804            profile_data: None,
805        };
806
807        optimizer.register_function(func1);
808        assert_eq!(optimizer.function_registry.len(), 1);
809    }
810
811    #[test]
812    fn test_optimization_stats() {
813        let config = OptimizationConfig::default();
814        let optimizer = FunctionOptimizer::new(config);
815
816        let stats = optimizer.get_optimization_stats();
817        assert_eq!(stats.total_optimizations, 0);
818        assert_eq!(stats.registered_functions, 0);
819        assert!(stats.active_patterns > 0);
820    }
821}