Skip to main content

scirs2_core/array_protocol/
grad.rs

1// Copyright (c) 2025, `SciRS2` Team
2//
3// Licensed under the Apache License, Version 2.0
4// (LICENSE-APACHE or http://www.apache.org/licenses/LICENSE-2.0)
5//
6
7//! Gradient computation support for the array protocol.
8//!
9//! This module provides automatic differentiation capabilities for arrays
10//! using the array protocol. It enables gradient computation for any array
11//! type that implements the `ArrayProtocol` trait.
12
13use crate::ndarray::compat::ArrayStatCompat;
14use ::ndarray::{Array, ArrayD, Dimension, Ix1, Ix2, IxDyn};
15
16use std::cell::RefCell;
17use std::collections::{HashMap, HashSet};
18use std::rc::Rc;
19
20use crate::array_protocol::operations::matmul;
21use crate::array_protocol::{ArrayProtocol, NdarrayWrapper};
22use crate::error::{CoreError, CoreResult, ErrorContext};
23
24/// Dictionary for storing parameter gradients
25#[derive(Clone)]
26pub struct GradientDict {
27    gradients: HashMap<String, Box<dyn ArrayProtocol>>,
28}
29
30impl std::fmt::Debug for GradientDict {
31    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
32        f.debug_struct("GradientDict")
33            .field(
34                "gradients",
35                &format!("{{keys: {:?}}}", self.gradients.keys().collect::<Vec<_>>()),
36            )
37            .finish()
38    }
39}
40
41impl GradientDict {
42    /// Create a new empty gradient dictionary
43    pub fn new() -> Self {
44        Self {
45            gradients: HashMap::new(),
46        }
47    }
48
49    /// Insert a gradient for a parameter
50    pub fn insert(&mut self, name: String, gradient: Box<dyn ArrayProtocol>) {
51        self.gradients.insert(name, gradient);
52    }
53
54    /// Get a gradient by parameter name
55    pub fn get(&self, name: &str) -> Option<&dyn ArrayProtocol> {
56        self.gradients.get(name).map(|b| b.as_ref())
57    }
58
59    /// Get a mutable reference to a gradient by parameter name
60    pub fn get_mut(&mut self, name: &str) -> Option<&mut Box<dyn ArrayProtocol>> {
61        self.gradients.get_mut(name)
62    }
63
64    /// Iterate over parameter names and gradients
65    pub fn iter(&self) -> impl Iterator<Item = (&String, &Box<dyn ArrayProtocol>)> {
66        self.gradients.iter()
67    }
68
69    /// Merge another gradient dictionary into this one
70    pub fn merge(&mut self, other: GradientDict) {
71        for (name, gradient) in other.gradients {
72            self.gradients.insert(name, gradient);
73        }
74    }
75
76    /// Check if the dictionary is empty
77    pub fn is_empty(&self) -> bool {
78        self.gradients.is_empty()
79    }
80
81    /// Get the number of gradients
82    pub fn len(&self) -> usize {
83        self.gradients.len()
84    }
85
86    /// Clear all gradients
87    pub fn clear(&mut self) {
88        self.gradients.clear();
89    }
90
91    /// Get all parameter names
92    pub fn keys(&self) -> impl Iterator<Item = &String> {
93        self.gradients.keys()
94    }
95
96    /// Get all gradients
97    pub fn values(&self) -> impl Iterator<Item = &Box<dyn ArrayProtocol>> {
98        self.gradients.values()
99    }
100}
101
102impl Default for GradientDict {
103    fn default() -> Self {
104        Self::new()
105    }
106}
107
108// Convert Box<dyn ArrayProtocol> to Rc<dyn ArrayProtocol> with proper trait object handling
109#[allow(dead_code)]
110fn boxed_to_rc(boxed: Box<dyn ArrayProtocol>) -> Rc<dyn ArrayProtocol> {
111    // `Rc<dyn Trait>` has a direct `From<Box<dyn Trait>>` impl in std: it
112    // reallocates the box's contents behind an `Rc` while preserving the exact
113    // concrete type and vtable. This previously downcast-enumerated only
114    // `NdarrayWrapper<f64, IxDyn>` and silently substituted a fabricated 1x1
115    // zero placeholder for every other type or dimension (e.g. the very common
116    // `NdarrayWrapper<f64, Ix2>` produced by `add`/`multiply` on `Array2`
117    // inputs) — a wrong-but-`Ok` result rather than a visible error.
118    Rc::from(boxed)
119}
120
121// Helper function to convert Box<dyn ArrayProtocol> to Rc<dyn ArrayProtocol>
122#[allow(dead_code)]
123fn box_to_rc_array_protocol(boxed: Box<dyn ArrayProtocol>) -> Rc<dyn ArrayProtocol> {
124    boxed_to_rc(boxed)
125}
126
127// Import the functions with the correct return type for our use
128#[allow(dead_code)]
129fn add(a: &dyn ArrayProtocol, b: &dyn ArrayProtocol) -> CoreResult<Box<dyn ArrayProtocol>> {
130    crate::array_protocol::operations::add(a, b).map_err(|e| e.into())
131}
132
133#[allow(dead_code)]
134fn multiply(a: &dyn ArrayProtocol, b: &dyn ArrayProtocol) -> CoreResult<Box<dyn ArrayProtocol>> {
135    crate::array_protocol::operations::multiply(a, b).map_err(|e| e.into())
136}
137
138#[allow(dead_code)]
139fn subtract(a: &dyn ArrayProtocol, b: &dyn ArrayProtocol) -> CoreResult<Box<dyn ArrayProtocol>> {
140    crate::array_protocol::operations::subtract(a, b).map_err(|e| e.into())
141}
142
143#[allow(dead_code)]
144fn ones_like(a: &dyn ArrayProtocol) -> CoreResult<Box<dyn ArrayProtocol>> {
145    // Create an array of ones with the same shape as the input
146    // Try different numeric types
147    if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<f64, IxDyn>>() {
148        let shape = a_array.as_array().shape();
149        let ones = ArrayD::<f64>::ones(IxDyn(shape));
150        Ok(Box::new(NdarrayWrapper::new(ones)) as Box<dyn ArrayProtocol>)
151    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<f32, IxDyn>>() {
152        let shape = a_array.as_array().shape();
153        let ones = ArrayD::<f32>::ones(IxDyn(shape));
154        Ok(Box::new(NdarrayWrapper::new(ones)) as Box<dyn ArrayProtocol>)
155    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<i32, IxDyn>>() {
156        let shape = a_array.as_array().shape();
157        let ones = ArrayD::<i32>::ones(IxDyn(shape));
158        Ok(Box::new(NdarrayWrapper::new(ones)) as Box<dyn ArrayProtocol>)
159    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<i64, IxDyn>>() {
160        let shape = a_array.as_array().shape();
161        let ones = ArrayD::<i64>::ones(IxDyn(shape));
162        Ok(Box::new(NdarrayWrapper::new(ones)) as Box<dyn ArrayProtocol>)
163    } else {
164        // Try to get shape information using the array protocol methods
165        let shape = a.shape().to_vec();
166        let ones = ArrayD::<f64>::ones(IxDyn(&shape));
167        Ok(Box::new(NdarrayWrapper::new(ones)) as Box<dyn ArrayProtocol>)
168    }
169}
170
171#[allow(dead_code)]
172fn broadcast_to(a: &dyn ArrayProtocol, shape: &[usize]) -> CoreResult<Box<dyn ArrayProtocol>> {
173    // Broadcast the array to the given shape
174    if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<f64, IxDyn>>() {
175        let array = a_array.as_array();
176        // Handle scalar broadcasting
177        if array.len() == 1 {
178            let value = array.iter().next().cloned().unwrap_or(0.0);
179            let broadcasted = ArrayD::<f64>::from_elem(IxDyn(shape), value);
180            Ok(Box::new(NdarrayWrapper::new(broadcasted)) as Box<dyn ArrayProtocol>)
181        } else if array.shape() == shape {
182            // Shapes already match
183            Ok(Box::new(NdarrayWrapper::new(array.clone())) as Box<dyn ArrayProtocol>)
184        } else {
185            // Implement basic broadcasting rules
186            let inputshape = array.shape();
187            let _ndim_diff = shape.len().saturating_sub(inputshape.len());
188
189            // Check if broadcasting is possible
190            let mut can_broadcast = true;
191            for i in 0..inputshape.len() {
192                let input_dim = inputshape[inputshape.len() - 1 - i];
193                let target_dim = shape[shape.len() - 1 - i];
194                if input_dim != 1 && input_dim != target_dim {
195                    can_broadcast = false;
196                    break;
197                }
198            }
199
200            if can_broadcast {
201                // Perform broadcasting by repeating data
202                if let Some(broadcasted_view) = array.broadcast(IxDyn(shape)) {
203                    let broadcasted = broadcasted_view.to_owned();
204                    Ok(Box::new(NdarrayWrapper::new(broadcasted)) as Box<dyn ArrayProtocol>)
205                } else {
206                    Err(CoreError::NotImplementedError(ErrorContext::new(
207                        "Broadcasting failed for these shapes".to_string(),
208                    )))
209                }
210            } else {
211                Err(CoreError::NotImplementedError(ErrorContext::new(
212                    "Incompatible shapes for broadcasting".to_string(),
213                )))
214            }
215        }
216    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<f32, IxDyn>>() {
217        let array = a_array.as_array();
218        if array.len() == 1 {
219            let value = array.iter().next().cloned().unwrap_or(0.0);
220            let broadcasted = ArrayD::<f32>::from_elem(IxDyn(shape), value);
221            Ok(Box::new(NdarrayWrapper::new(broadcasted)) as Box<dyn ArrayProtocol>)
222        } else if array.shape() == shape {
223            Ok(Box::new(NdarrayWrapper::new(array.clone())) as Box<dyn ArrayProtocol>)
224        } else if let Some(broadcasted_view) = array.broadcast(IxDyn(shape)) {
225            let broadcasted = broadcasted_view.to_owned();
226            Ok(Box::new(NdarrayWrapper::new(broadcasted)) as Box<dyn ArrayProtocol>)
227        } else {
228            Err(CoreError::NotImplementedError(ErrorContext::new(
229                "Broadcasting failed for these shapes".to_string(),
230            )))
231        }
232    } else {
233        // Fallback: create an array of ones with the target shape
234        let ones = ArrayD::<f64>::ones(IxDyn(shape));
235        Ok(Box::new(NdarrayWrapper::new(ones)) as Box<dyn ArrayProtocol>)
236    }
237}
238
239/// A node in the computation graph.
240#[derive(Clone)]
241struct Node {
242    /// The array value at this node.
243    value: Rc<dyn ArrayProtocol>,
244
245    /// Gradient with respect to the output.
246    grad: Option<Rc<dyn ArrayProtocol>>,
247
248    /// Operation that created this node.
249    op: Option<String>,
250
251    /// Input nodes for the operation that created this node.
252    inputs: Vec<GradientTensor>,
253
254    /// Whether gradient computation is required for this node.
255    requiresgrad: bool,
256
257    /// Whether this node is a leaf node (parameter or input).
258    is_leaf: bool,
259}
260
261impl Node {
262    /// Create a new leaf node.
263    fn leaf(requiresgrad: bool) -> Self {
264        Self {
265            value: Rc::new(NdarrayWrapper::new(
266                crate::ndarray::Array0::<f64>::zeros(()),
267            )) as Rc<dyn ArrayProtocol>,
268            grad: None,
269            op: None,
270            inputs: Vec::new(),
271            requiresgrad,
272            is_leaf: true,
273        }
274    }
275
276    /// Create a new operation node.
277    fn new_op(value: Rc<dyn ArrayProtocol>, op: String, inputs: Vec<GradientTensor>) -> Self {
278        let requiresgrad = inputs.iter().any(|x| x.requiresgrad());
279
280        Self {
281            value,
282            grad: None,
283            op: Some(op),
284            inputs,
285            requiresgrad,
286            is_leaf: false,
287        }
288    }
289}
290
291/// Tensor with gradient tracking capabilities.
292#[derive(Clone)]
293pub struct GradientTensor {
294    /// The node in the computation graph.
295    node: Rc<RefCell<Node>>,
296}
297
298impl GradientTensor {
299    /// Create a new gradient tensor from a value.
300    pub fn new(value: Rc<dyn ArrayProtocol>, requiresgrad: bool) -> Self {
301        let mut node_inner = Node::leaf(requiresgrad);
302        node_inner.value = value;
303        node_inner.grad = None;
304        let node = Rc::new(RefCell::new(node_inner));
305        Self { node }
306    }
307
308    /// Create a new gradient tensor from an array.
309    pub fn from_array<T, D>(array: Array<T, D>, requiresgrad: bool) -> Self
310    where
311        T: Clone + Send + Sync + 'static,
312        D: Dimension + Send + Sync + 'static,
313    {
314        let value = Rc::new(NdarrayWrapper::new(array)) as Rc<dyn ArrayProtocol>;
315        Self::new(value, requiresgrad)
316    }
317
318    /// Get the value of the tensor.
319    pub fn value(&self) -> Rc<dyn ArrayProtocol> {
320        self.node.borrow().value.clone()
321    }
322
323    /// Get the gradient of the tensor.
324    pub fn grad_2(&self) -> Option<Rc<dyn ArrayProtocol>> {
325        self.node.borrow().grad.clone()
326    }
327
328    /// Check if gradient computation is required for this tensor.
329    pub fn requiresgrad(&self) -> bool {
330        self.node.borrow().requiresgrad
331    }
332
333    /// Set whether gradient computation is required for this tensor.
334    pub fn set_requiresgrad(&mut self, requiresgrad: bool) {
335        self.node.borrow_mut().requiresgrad = requiresgrad;
336    }
337
338    /// Check if this tensor is a leaf node.
339    pub fn is_leaf(&self) -> bool {
340        self.node.borrow().is_leaf
341    }
342
343    /// Create a new tensor from an operation.
344    fn from_op(value: Rc<dyn ArrayProtocol>, op: String, inputs: Vec<GradientTensor>) -> Self {
345        let node = Rc::new(RefCell::new(Node::new_op(value, op, inputs)));
346        Self { node }
347    }
348
349    /// Set the value of the tensor (for updating variables during optimization).
350    pub fn set_value(&mut self, newvalue: Rc<dyn ArrayProtocol>) {
351        self.node.borrow_mut().grad = None; // Clear gradient when value changes
352        self.node.borrow_mut().value = newvalue;
353    }
354
355    /// Backward pass to compute gradients.
356    pub fn backward(&self) -> CoreResult<()> {
357        // Initialize gradient as ones with the same shape as value.
358        // `ArrayProtocol::shape()` works for any concrete dimension type,
359        // unlike downcasting to a specific `NdarrayWrapper<f64, D>` — this
360        // used to only recognize `IxDyn` and silently fall back to a
361        // scalar-shaped gradient (`[1]`) for any other dimension (e.g. the
362        // very common `Ix2`), corrupting the shape the entire backward pass
363        // is seeded with.
364        let gradshape = IxDyn(self.value().shape());
365
366        let grad_array = Array::<f64, IxDyn>::ones(gradshape);
367        let grad = Rc::new(NdarrayWrapper::new(grad_array)) as Rc<dyn ArrayProtocol>;
368
369        // Perform backward pass
370        self.backward_with_grad(grad)
371    }
372
373    /// Backward pass with a specific gradient.
374    fn backward_with_grad(&self, grad: Rc<dyn ArrayProtocol>) -> CoreResult<()> {
375        // Set the gradient of this tensor
376        self.node.borrow_mut().grad = Some(grad.clone());
377
378        // Create a topologically sorted list of nodes
379        let mut visited = HashSet::new();
380        let mut topo = Vec::new();
381
382        // Helper function for topological sort
383        fn build_topo(
384            tensor: &GradientTensor,
385            visited: &mut HashSet<*const RefCell<Node>>,
386            topo: &mut Vec<GradientTensor>,
387        ) {
388            let node_ptr = Rc::as_ptr(&tensor.node);
389            if !visited.contains(&node_ptr) {
390                visited.insert(node_ptr);
391
392                // Visit all inputs first
393                for input in &tensor.node.borrow().inputs {
394                    build_topo(input, visited, topo);
395                }
396
397                // Then add this node
398                topo.push(tensor.clone());
399            }
400        }
401
402        // Build topological sort
403        build_topo(self, &mut visited, &mut topo);
404
405        // Perform backward pass in reverse topological order
406        for node in topo.iter().rev() {
407            // Only compute gradients for nodes that require it
408            if !node.requiresgrad() {
409                continue;
410            }
411
412            // Get the gradient of this node
413            let node_grad = match node.grad_2() {
414                Some(g) => g,
415                None => continue, // Skip nodes with no gradient
416            };
417
418            // If this is a leaf node, we're done
419            if node.is_leaf() {
420                continue;
421            }
422
423            // Get the operation and inputs
424            let op = match &node.node.borrow().op {
425                Some(op) => op.clone(),
426                None => continue, // Skip nodes with no operation
427            };
428
429            let inputs = node.node.borrow().inputs.clone();
430
431            // Compute gradients for inputs based on the operation
432            match op.as_str() {
433                "add" => {
434                    // For addition, gradient flows directly to both inputs
435                    for input in &inputs {
436                        if input.requiresgrad() {
437                            let mut input_node = input.node.borrow_mut();
438                            if let Some(input_grad) = &input_node.grad {
439                                // Accumulate gradients
440                                if let Ok(sum) = add(input_grad.as_ref(), node_grad.as_ref()) {
441                                    input_node.grad = Some(sum.into());
442                                }
443                            } else {
444                                input_node.grad = Some(node_grad.clone());
445                            }
446                        }
447                    }
448                }
449                "multiply"
450                    // For element-wise multiplication, input_grad = output_grad * other_input
451                    if inputs.len() == 2 => {
452                        let (a, b) = (&inputs[0], &inputs[1]);
453
454                        // Compute grad_a = grad_out * b
455                        if a.requiresgrad() {
456                            let b_value = b.value();
457                            if let Ok(grad_a) = multiply(node_grad.as_ref(), b_value.as_ref()) {
458                                let mut a_node = a.node.borrow_mut();
459                                if let Some(a_grad) = &a_node.grad {
460                                    // Accumulate gradients
461                                    if let Ok(sum) = add(a_grad.as_ref(), grad_a.as_ref()) {
462                                        a_node.grad = Some(box_to_rc_array_protocol(sum));
463                                    }
464                                } else {
465                                    a_node.grad = Some(box_to_rc_array_protocol(grad_a));
466                                }
467                            }
468                        }
469
470                        // Compute grad_b = grad_out * a
471                        if b.requiresgrad() {
472                            let a_value = a.value();
473                            if let Ok(grad_b) = multiply(node_grad.as_ref(), a_value.as_ref()) {
474                                let mut b_node = b.node.borrow_mut();
475                                if let Some(b_grad) = &b_node.grad {
476                                    // Accumulate gradients
477                                    if let Ok(sum) = add(b_grad.as_ref(), grad_b.as_ref()) {
478                                        b_node.grad = Some(box_to_rc_array_protocol(sum));
479                                    }
480                                } else {
481                                    b_node.grad = Some(box_to_rc_array_protocol(grad_b));
482                                }
483                            }
484                        }
485                    }
486                "matmul"
487                    // For matrix multiplication, the gradients are more complex:
488                    // grad_a = grad_out @ b.T
489                    // grad_b = a.T @ grad_out
490                    if inputs.len() == 2 => {
491                        let (a, b) = (&inputs[0], &inputs[1]);
492
493                        // Compute grad_a = grad_out @ b.T
494                        if a.requiresgrad() {
495                            if let (Some(b_array), Some(grad_out_array)) = (
496                                b.value()
497                                    .as_any()
498                                    .downcast_ref::<NdarrayWrapper<f64, IxDyn>>(),
499                                node_grad
500                                    .as_any()
501                                    .downcast_ref::<NdarrayWrapper<f64, IxDyn>>(),
502                            ) {
503                                let b_array_val = b_array.as_array();
504                                let grad_out_array_val = grad_out_array.as_array();
505
506                                // Transpose b: b_t = b.t()
507                                let b_t = b_array_val.t();
508
509                                // Matrix multiplication: grad_a = grad_out @ b_t
510                                // Convert to Array2 for more deterministic dot behavior
511                                let grad_outshape = grad_out_array_val.shape();
512                                let grad_out_rows = grad_outshape[0];
513                                let grad_out_cols = if grad_outshape.len() > 1 {
514                                    grad_outshape.iter().skip(1).product()
515                                } else {
516                                    1
517                                };
518                                let grad_out_2d = grad_out_array_val
519                                    .clone()
520                                    .into_shape_with_order((grad_out_rows, grad_out_cols))
521                                    .expect("Operation failed");
522
523                                let b_tshape = b_t.shape();
524                                let b_t_rows = b_tshape[0];
525                                let b_t_cols = if b_tshape.len() > 1 {
526                                    b_tshape.iter().skip(1).product()
527                                } else {
528                                    1
529                                };
530                                let b_t_2d = b_t
531                                    .clone()
532                                    .into_shape_with_order((b_t_rows, b_t_cols))
533                                    .expect("Operation failed");
534
535                                // Now use dot with 2D arrays
536                                let grad_a_val = grad_out_2d.dot(&b_t_2d);
537
538                                // Convert back to IxDyn for consistency
539                                let grad_a_dyn = grad_a_val.into_dyn();
540                                let grad_a = NdarrayWrapper::new(grad_a_dyn);
541
542                                // Update a's gradient
543                                let mut a_node = a.node.borrow_mut();
544                                if let Some(a_grad) = &a_node.grad {
545                                    if let (Some(a_grad_array), Some(grad_a_array)) = (
546                                        a_grad
547                                            .as_any()
548                                            .downcast_ref::<NdarrayWrapper<f64, IxDyn>>(),
549                                        grad_a
550                                            .as_any()
551                                            .downcast_ref::<NdarrayWrapper<f64, IxDyn>>(),
552                                    ) {
553                                        // Accumulate gradients: a_grad += grad_a
554                                        let sum = a_grad_array.as_array() + grad_a_array.as_array();
555                                        a_node.grad = Some(Rc::new(NdarrayWrapper::new(sum)));
556                                    }
557                                } else {
558                                    // Use Box<dyn ArrayProtocol> and convert to Rc
559                                    a_node.grad = Some(Rc::new(grad_a));
560                                }
561                            }
562                        }
563
564                        // Compute grad_b = a.T @ grad_out
565                        if b.requiresgrad() {
566                            if let (Some(a_array), Some(grad_out_array)) = (
567                                a.value()
568                                    .as_any()
569                                    .downcast_ref::<NdarrayWrapper<f64, IxDyn>>(),
570                                node_grad
571                                    .as_any()
572                                    .downcast_ref::<NdarrayWrapper<f64, IxDyn>>(),
573                            ) {
574                                let a_array_val = a_array.as_array();
575                                let grad_out_array_val = grad_out_array.as_array();
576
577                                // Transpose a: a_t = a.t()
578                                let a_t = a_array_val.t();
579
580                                // Matrix multiplication: grad_b = a_t @ grad_out
581                                // Convert to Array2 for more deterministic dot behavior
582                                let grad_outshape = grad_out_array_val.shape();
583                                let grad_out_rows = grad_outshape[0];
584                                let grad_out_cols = if grad_outshape.len() > 1 {
585                                    grad_outshape.iter().skip(1).product()
586                                } else {
587                                    1
588                                };
589                                let grad_out_2d = grad_out_array_val
590                                    .clone()
591                                    .into_shape_with_order((grad_out_rows, grad_out_cols))
592                                    .expect("Operation failed");
593
594                                let a_tshape = a_t.shape();
595                                let a_t_rows = a_tshape[0];
596                                let a_t_cols = if a_tshape.len() > 1 {
597                                    a_tshape.iter().skip(1).product()
598                                } else {
599                                    1
600                                };
601                                let a_t_2d = a_t
602                                    .clone()
603                                    .into_shape_with_order((a_t_rows, a_t_cols))
604                                    .expect("Operation failed");
605
606                                // Now use dot with 2D arrays
607                                let grad_b_val = a_t_2d.dot(&grad_out_2d);
608
609                                // Convert back to IxDyn for consistency
610                                let grad_b_dyn = grad_b_val.into_dyn();
611                                let grad_b = NdarrayWrapper::new(grad_b_dyn);
612
613                                // Update b's gradient
614                                let mut b_node = b.node.borrow_mut();
615                                if let Some(b_grad) = &b_node.grad {
616                                    if let (Some(b_grad_array), Some(grad_b_array)) = (
617                                        b_grad
618                                            .as_any()
619                                            .downcast_ref::<NdarrayWrapper<f64, IxDyn>>(),
620                                        grad_b
621                                            .as_any()
622                                            .downcast_ref::<NdarrayWrapper<f64, IxDyn>>(),
623                                    ) {
624                                        // Accumulate gradients: b_grad += grad_b
625                                        let sum = b_grad_array.as_array() + grad_b_array.as_array();
626                                        b_node.grad = Some(Rc::new(NdarrayWrapper::new(sum)));
627                                    }
628                                } else {
629                                    // Use Box<dyn ArrayProtocol> and convert to Rc
630                                    b_node.grad = Some(Rc::new(grad_b));
631                                }
632                            }
633                        }
634                    }
635                "subtract"
636                    // For subtraction: a - b, grad_a = grad_out, grad_b = -grad_out
637                    if inputs.len() == 2 => {
638                        let (a, b) = (&inputs[0], &inputs[1]);
639
640                        // Compute grad_a = grad_out
641                        if a.requiresgrad() {
642                            let mut a_node = a.node.borrow_mut();
643                            if let Some(a_grad) = &a_node.grad {
644                                // Accumulate gradients
645                                if let Ok(sum) = add(a_grad.as_ref(), node_grad.as_ref()) {
646                                    a_node.grad = Some(box_to_rc_array_protocol(sum));
647                                }
648                            } else {
649                                a_node.grad = Some(node_grad.clone());
650                            }
651                        }
652
653                        // Compute grad_b = -grad_out
654                        if b.requiresgrad() {
655                            if let Ok(neg_grad) = multiply_by_scalar(node_grad.as_ref(), -1.0) {
656                                let mut b_node = b.node.borrow_mut();
657                                if let Some(b_grad) = &b_node.grad {
658                                    // Accumulate gradients
659                                    if let Ok(sum) = add(b_grad.as_ref(), neg_grad.as_ref()) {
660                                        b_node.grad = Some(box_to_rc_array_protocol(sum));
661                                    }
662                                } else {
663                                    b_node.grad = Some(box_to_rc_array_protocol(neg_grad));
664                                }
665                            }
666                        }
667                    }
668                "divide"
669                    // For division: a / b, grad_a = grad_out / b, grad_b = -grad_out * a / b^2
670                    if inputs.len() == 2 => {
671                        let (a, b) = (&inputs[0], &inputs[1]);
672
673                        // Compute grad_a = grad_out / b
674                        if a.requiresgrad() {
675                            let b_value = b.value();
676                            if let Ok(grad_a) = divide(node_grad.as_ref(), b_value.as_ref()) {
677                                let mut a_node = a.node.borrow_mut();
678                                if let Some(a_grad) = &a_node.grad {
679                                    // Accumulate gradients
680                                    if let Ok(sum) = add(a_grad.as_ref(), grad_a.as_ref()) {
681                                        a_node.grad = Some(box_to_rc_array_protocol(sum));
682                                    }
683                                } else {
684                                    a_node.grad = Some(box_to_rc_array_protocol(grad_a));
685                                }
686                            }
687                        }
688
689                        // Compute grad_b = -grad_out * a / b^2
690                        if b.requiresgrad() {
691                            let a_value = a.value();
692                            let b_value = b.value();
693
694                            // Compute b^2
695                            if let Ok(b_squared) = multiply(b_value.as_ref(), b_value.as_ref()) {
696                                // Compute grad_out * a
697                                if let Ok(grad_times_a) =
698                                    multiply(node_grad.as_ref(), a_value.as_ref())
699                                {
700                                    // Compute grad_out * a / b^2
701                                    if let Ok(div_result) =
702                                        divide(grad_times_a.as_ref(), b_squared.as_ref())
703                                    {
704                                        // Negate: -grad_out * a / b^2
705                                        if let Ok(grad_b) =
706                                            multiply_by_scalar(div_result.as_ref(), -1.0)
707                                        {
708                                            let mut b_node = b.node.borrow_mut();
709                                            if let Some(b_grad) = &b_node.grad {
710                                                // Accumulate gradients
711                                                if let Ok(sum) =
712                                                    add(b_grad.as_ref(), grad_b.as_ref())
713                                                {
714                                                    b_node.grad =
715                                                        Some(box_to_rc_array_protocol(sum));
716                                                }
717                                            } else {
718                                                b_node.grad =
719                                                    Some(box_to_rc_array_protocol(grad_b));
720                                            }
721                                        }
722                                    }
723                                }
724                            }
725                        }
726                    }
727                "sigmoid"
728                    // For sigmoid: grad_input = grad_out * sigmoid * (1 - sigmoid)
729                    if inputs.len() == 1 => {
730                        let input = &inputs[0];
731
732                        if input.requiresgrad() {
733                            // Get the output value (sigmoid result)
734                            let sigmoid_value = node.value();
735
736                            // Compute 1 - sigmoid
737                            if let Ok(ones) = ones_like(sigmoid_value.as_ref()) {
738                                if let Ok(one_minus_sigmoid) =
739                                    subtract(ones.as_ref(), sigmoid_value.as_ref())
740                                {
741                                    // Compute sigmoid * (1 - sigmoid)
742                                    if let Ok(sigmoid_deriv) =
743                                        multiply(sigmoid_value.as_ref(), one_minus_sigmoid.as_ref())
744                                    {
745                                        // Compute grad_out * sigmoid * (1 - sigmoid)
746                                        if let Ok(grad_input) =
747                                            multiply(node_grad.as_ref(), sigmoid_deriv.as_ref())
748                                        {
749                                            let mut input_node = input.node.borrow_mut();
750                                            if let Some(input_grad) = &input_node.grad {
751                                                // Accumulate gradients
752                                                if let Ok(sum) =
753                                                    add(input_grad.as_ref(), grad_input.as_ref())
754                                                {
755                                                    input_node.grad =
756                                                        Some(box_to_rc_array_protocol(sum));
757                                                }
758                                            } else {
759                                                input_node.grad =
760                                                    Some(box_to_rc_array_protocol(grad_input));
761                                            }
762                                        }
763                                    }
764                                }
765                            }
766                        }
767                    }
768                "mean"
769                    // For mean: grad_input = grad_out / n (where n is the number of elements)
770                    if inputs.len() == 1 => {
771                        let input = &inputs[0];
772
773                        if input.requiresgrad() {
774                            // Get the number of elements
775                            let input_value = input.value();
776                            if let Some(inputarray) = input_value
777                                .as_any()
778                                .downcast_ref::<NdarrayWrapper<f64, IxDyn>>()
779                            {
780                                let n_elements = inputarray.as_array().len() as f64;
781
782                                // Compute grad_input = grad_out / n
783                                if let Ok(grad_input) =
784                                    multiply_by_scalar(node_grad.as_ref(), 1.0 / n_elements)
785                                {
786                                    // Broadcast the gradient to match input shape
787                                    if let Ok(broadcasted_grad) = broadcast_to(
788                                        grad_input.as_ref(),
789                                        inputarray.as_array().shape(),
790                                    ) {
791                                        let mut input_node = input.node.borrow_mut();
792                                        if let Some(input_grad) = &input_node.grad {
793                                            // Accumulate gradients
794                                            if let Ok(sum) =
795                                                add(input_grad.as_ref(), broadcasted_grad.as_ref())
796                                            {
797                                                input_node.grad =
798                                                    Some(box_to_rc_array_protocol(sum));
799                                            }
800                                        } else {
801                                            input_node.grad =
802                                                Some(box_to_rc_array_protocol(broadcasted_grad));
803                                        }
804                                    }
805                                }
806                            }
807                        }
808                    }
809                _ => {
810                    // Other operations would be implemented here
811                }
812            }
813        }
814
815        Ok(())
816    }
817
818    /// Detach this tensor from the computation graph.
819    pub fn detach(&self) -> Self {
820        GradientTensor::new(self.value(), false)
821    }
822}
823
824/// Implementations of gradient-aware operations
825/// Addition operation with gradient tracking.
826#[allow(dead_code)]
827pub fn grad_add(a: &GradientTensor, b: &GradientTensor) -> CoreResult<GradientTensor> {
828    let a_value = a.value();
829    let b_value = b.value();
830
831    // Perform addition
832    let result = add(a_value.as_ref(), b_value.as_ref())?;
833
834    // Create a new gradient tensor - explicitly convert Box to Rc
835    let result_rc: Rc<dyn ArrayProtocol> = box_to_rc_array_protocol(result);
836    Ok(GradientTensor::from_op(
837        result_rc,
838        "add".to_string(),
839        vec![a.clone(), b.clone()],
840    ))
841}
842
843/// Element-wise multiplication with gradient tracking.
844#[allow(dead_code)]
845pub fn grad_multiply(a: &GradientTensor, b: &GradientTensor) -> CoreResult<GradientTensor> {
846    let a_value = a.value();
847    let b_value = b.value();
848
849    // Perform multiplication
850    let result = multiply(a_value.as_ref(), b_value.as_ref())?;
851
852    // Create a new gradient tensor - explicitly convert Box to Rc
853    let result_rc: Rc<dyn ArrayProtocol> = box_to_rc_array_protocol(result);
854    Ok(GradientTensor::from_op(
855        result_rc,
856        "multiply".to_string(),
857        vec![a.clone(), b.clone()],
858    ))
859}
860
861/// Matrix multiplication with gradient tracking.
862#[allow(dead_code)]
863pub fn grad_matmul(a: &GradientTensor, b: &GradientTensor) -> CoreResult<GradientTensor> {
864    let a_value = a.value();
865    let b_value = b.value();
866
867    // Perform matrix multiplication
868    let result = matmul(a_value.as_ref(), b_value.as_ref())?;
869
870    // Create a new gradient tensor - explicitly convert Box to Rc
871    let result_rc: Rc<dyn ArrayProtocol> = box_to_rc_array_protocol(result);
872    Ok(GradientTensor::from_op(
873        result_rc,
874        "matmul".to_string(),
875        vec![a.clone(), b.clone()],
876    ))
877}
878
879/// Subtraction with gradient tracking.
880#[allow(dead_code)]
881pub fn grad_subtract(a: &GradientTensor, b: &GradientTensor) -> CoreResult<GradientTensor> {
882    let a_value = a.value();
883    let b_value = b.value();
884
885    // Perform subtraction
886    let result = subtract(a_value.as_ref(), b_value.as_ref())?;
887
888    // Create a new gradient tensor - explicitly convert Box to Rc
889    let result_rc: Rc<dyn ArrayProtocol> = box_to_rc_array_protocol(result);
890    Ok(GradientTensor::from_op(
891        result_rc,
892        "subtract".to_string(),
893        vec![a.clone(), b.clone()],
894    ))
895}
896
897/// Division with gradient tracking.
898#[allow(dead_code)]
899pub fn grad_divide(a: &GradientTensor, b: &GradientTensor) -> CoreResult<GradientTensor> {
900    let a_value = a.value();
901    let b_value = b.value();
902
903    // Perform division
904    let result = divide(a_value.as_ref(), b_value.as_ref())?;
905
906    // Create a new gradient tensor - explicitly convert Box to Rc
907    let result_rc: Rc<dyn ArrayProtocol> = box_to_rc_array_protocol(result);
908    Ok(GradientTensor::from_op(
909        result_rc,
910        "divide".to_string(),
911        vec![a.clone(), b.clone()],
912    ))
913}
914
915/// Sigmoid activation with gradient tracking.
916#[allow(dead_code)]
917pub fn grad_sigmoid(a: &GradientTensor) -> CoreResult<GradientTensor> {
918    let a_value = a.value();
919
920    // Perform sigmoid: 1 / (1 + exp(-x))
921    if let Some(a_array) = a_value
922        .as_any()
923        .downcast_ref::<NdarrayWrapper<f64, IxDyn>>()
924    {
925        let array = a_array.as_array();
926        let result = array.mapv(|x| 1.0 / (1.0 + (-x).exp()));
927        let result_wrapped = NdarrayWrapper::new(result);
928        let result_rc: Rc<dyn ArrayProtocol> = Rc::new(result_wrapped);
929        Ok(GradientTensor::from_op(
930            result_rc,
931            "sigmoid".to_string(),
932            vec![a.clone()],
933        ))
934    } else if let Some(a_array) = a_value
935        .as_any()
936        .downcast_ref::<NdarrayWrapper<f32, IxDyn>>()
937    {
938        let array = a_array.as_array();
939        let result = array.mapv(|x| 1.0f32 / (1.0f32 + (-x).exp()));
940        let result_wrapped = NdarrayWrapper::new(result);
941        let result_rc: Rc<dyn ArrayProtocol> = Rc::new(result_wrapped);
942        Ok(GradientTensor::from_op(
943            result_rc,
944            "sigmoid".to_string(),
945            vec![a.clone()],
946        ))
947    } else {
948        Err(CoreError::NotImplementedError(ErrorContext::new(
949            "sigmoid not implemented for this array type".to_string(),
950        )))
951    }
952}
953
954/// Mean reduction with gradient tracking.
955#[allow(dead_code)]
956pub fn grad_mean(a: &GradientTensor) -> CoreResult<GradientTensor> {
957    let a_value = a.value();
958
959    // Perform mean reduction
960    if let Some(a_array) = a_value
961        .as_any()
962        .downcast_ref::<NdarrayWrapper<f64, IxDyn>>()
963    {
964        let array = a_array.as_array();
965        let mean_value = array.mean_or(0.0);
966        let result = ArrayD::<f64>::from_elem(IxDyn(&[1]), mean_value);
967        let result_wrapped = NdarrayWrapper::new(result);
968        let result_rc: Rc<dyn ArrayProtocol> = Rc::new(result_wrapped);
969        Ok(GradientTensor::from_op(
970            result_rc,
971            "mean".to_string(),
972            vec![a.clone()],
973        ))
974    } else if let Some(a_array) = a_value
975        .as_any()
976        .downcast_ref::<NdarrayWrapper<f32, IxDyn>>()
977    {
978        let array = a_array.as_array();
979        let mean_value = array.mean_or(0.0f32);
980        let result = ArrayD::<f32>::from_elem(IxDyn(&[1]), mean_value);
981        let result_wrapped = NdarrayWrapper::new(result);
982        let result_rc: Rc<dyn ArrayProtocol> = Rc::new(result_wrapped);
983        Ok(GradientTensor::from_op(
984            result_rc,
985            "mean".to_string(),
986            vec![a.clone()],
987        ))
988    } else if let Some(a_array) = a_value
989        .as_any()
990        .downcast_ref::<NdarrayWrapper<i32, IxDyn>>()
991    {
992        let array = a_array.as_array();
993        let mean_value = if array.is_empty() {
994            0.0f64
995        } else {
996            array.iter().map(|&x| x as f64).sum::<f64>() / array.len() as f64
997        };
998        let result = ArrayD::<f64>::from_elem(IxDyn(&[1]), mean_value);
999        let result_wrapped = NdarrayWrapper::new(result);
1000        let result_rc: Rc<dyn ArrayProtocol> = Rc::new(result_wrapped);
1001        Ok(GradientTensor::from_op(
1002            result_rc,
1003            "mean".to_string(),
1004            vec![a.clone()],
1005        ))
1006    } else if let Some(a_array) = a_value
1007        .as_any()
1008        .downcast_ref::<NdarrayWrapper<i64, IxDyn>>()
1009    {
1010        let array = a_array.as_array();
1011        let mean_value = if array.is_empty() {
1012            0.0f64
1013        } else {
1014            array.iter().map(|&x| x as f64).sum::<f64>() / array.len() as f64
1015        };
1016        let result = ArrayD::<f64>::from_elem(IxDyn(&[1]), mean_value);
1017        let result_wrapped = NdarrayWrapper::new(result);
1018        let result_rc: Rc<dyn ArrayProtocol> = Rc::new(result_wrapped);
1019        Ok(GradientTensor::from_op(
1020            result_rc,
1021            "mean".to_string(),
1022            vec![a.clone()],
1023        ))
1024    } else if let Some(a_array) = a_value.as_any().downcast_ref::<NdarrayWrapper<u8, IxDyn>>() {
1025        let array = a_array.as_array();
1026        let mean_value = if array.is_empty() {
1027            0.0f64
1028        } else {
1029            array.iter().map(|&x| x as f64).sum::<f64>() / array.len() as f64
1030        };
1031        let result = ArrayD::<f64>::from_elem(IxDyn(&[1]), mean_value);
1032        let result_wrapped = NdarrayWrapper::new(result);
1033        let result_rc: Rc<dyn ArrayProtocol> = Rc::new(result_wrapped);
1034        Ok(GradientTensor::from_op(
1035            result_rc,
1036            "mean".to_string(),
1037            vec![a.clone()],
1038        ))
1039    } else if let Some(a_array) = a_value
1040        .as_any()
1041        .downcast_ref::<NdarrayWrapper<u16, IxDyn>>()
1042    {
1043        let array = a_array.as_array();
1044        let mean_value = if array.is_empty() {
1045            0.0f64
1046        } else {
1047            array.iter().map(|&x| x as f64).sum::<f64>() / array.len() as f64
1048        };
1049        let result = ArrayD::<f64>::from_elem(IxDyn(&[1]), mean_value);
1050        let result_wrapped = NdarrayWrapper::new(result);
1051        let result_rc: Rc<dyn ArrayProtocol> = Rc::new(result_wrapped);
1052        Ok(GradientTensor::from_op(
1053            result_rc,
1054            "mean".to_string(),
1055            vec![a.clone()],
1056        ))
1057    } else if let Some(a_array) = a_value
1058        .as_any()
1059        .downcast_ref::<NdarrayWrapper<u32, IxDyn>>()
1060    {
1061        let array = a_array.as_array();
1062        let mean_value = if array.is_empty() {
1063            0.0f64
1064        } else {
1065            array.iter().map(|&x| x as f64).sum::<f64>() / array.len() as f64
1066        };
1067        let result = ArrayD::<f64>::from_elem(IxDyn(&[1]), mean_value);
1068        let result_wrapped = NdarrayWrapper::new(result);
1069        let result_rc: Rc<dyn ArrayProtocol> = Rc::new(result_wrapped);
1070        Ok(GradientTensor::from_op(
1071            result_rc,
1072            "mean".to_string(),
1073            vec![a.clone()],
1074        ))
1075    } else if let Some(a_array) = a_value
1076        .as_any()
1077        .downcast_ref::<NdarrayWrapper<u64, IxDyn>>()
1078    {
1079        let array = a_array.as_array();
1080        let mean_value = if array.is_empty() {
1081            0.0f64
1082        } else {
1083            array.iter().map(|&x| x as f64).sum::<f64>() / array.len() as f64
1084        };
1085        let result = ArrayD::<f64>::from_elem(IxDyn(&[1]), mean_value);
1086        let result_wrapped = NdarrayWrapper::new(result);
1087        let result_rc: Rc<dyn ArrayProtocol> = Rc::new(result_wrapped);
1088        Ok(GradientTensor::from_op(
1089            result_rc,
1090            "mean".to_string(),
1091            vec![a.clone()],
1092        ))
1093    } else {
1094        Err(CoreError::NotImplementedError(ErrorContext::new(
1095            "mean not implemented for this array type".to_string(),
1096        )))
1097    }
1098}
1099
1100/// Gradient-aware variable that can be optimized.
1101pub struct Variable {
1102    /// The gradient tensor.
1103    tensor: GradientTensor,
1104
1105    /// Name for the variable.
1106    name: String,
1107}
1108
1109impl Variable {
1110    /// Create a new variable from an array.
1111    pub fn new<T, D>(name: &str, array: Array<T, D>) -> Self
1112    where
1113        T: Clone + Send + Sync + 'static,
1114        D: Dimension + Send + Sync + 'static,
1115    {
1116        let tensor = GradientTensor::from_array(array, true);
1117        Self {
1118            tensor,
1119            name: name.to_string(),
1120        }
1121    }
1122
1123    /// Get the gradient tensor.
1124    pub const fn tensor(&self) -> &GradientTensor {
1125        &self.tensor
1126    }
1127
1128    /// Get the value of the variable.
1129    pub fn value(&self) -> Rc<dyn ArrayProtocol> {
1130        self.tensor.value()
1131    }
1132
1133    /// Get the gradient of the variable.
1134    pub fn grad_2(&self) -> Option<Rc<dyn ArrayProtocol>> {
1135        self.tensor.grad_2()
1136    }
1137
1138    /// Get the name of the variable.
1139    pub fn name(&self) -> &str {
1140        &self.name
1141    }
1142
1143    /// Set the gradient of the variable
1144    pub fn set_gradient(&mut self, gradient: Box<dyn ArrayProtocol>) -> CoreResult<()> {
1145        // Convert Box to Rc
1146        let gradient_rc = self.box_to_rc(gradient);
1147
1148        // Set gradient on the tensor node
1149        self.tensor.node.borrow_mut().grad = Some(gradient_rc);
1150        Ok(())
1151    }
1152
1153    /// Set the value of the variable (for updating during optimization).
1154    pub fn set_value(&mut self, newvalue: Box<dyn ArrayProtocol>) {
1155        let newvalue_rc = self.box_to_rc(newvalue);
1156        self.tensor.set_value(newvalue_rc);
1157    }
1158
1159    /// Helper to convert Box<dyn ArrayProtocol> to Rc<dyn ArrayProtocol>
1160    fn box_to_rc(&self, boxed: Box<dyn ArrayProtocol>) -> Rc<dyn ArrayProtocol> {
1161        // See the module-level `boxed_to_rc`: `Rc::from` preserves the exact
1162        // concrete type via std's `Rc<T>: From<Box<T>>` impl, rather than
1163        // downcast-enumerating only `NdarrayWrapper<f64, IxDyn>` and silently
1164        // substituting a fabricated 1x1 zero placeholder for anything else
1165        // (e.g. an update or gradient with a different dimension type).
1166        Rc::from(boxed)
1167    }
1168}
1169
1170/// Trait for optimizers that update variables.
1171pub trait Optimizer {
1172    /// Step the optimizer to update variables.
1173    fn step(&mut self) -> CoreResult<()>;
1174
1175    /// Zero all gradients.
1176    fn zero_grad(&mut self);
1177
1178    /// Add a variable to the optimizer.
1179    fn add_variable(&mut self, var: Variable);
1180
1181    /// Get all variables managed by the optimizer.
1182    fn variables(&self) -> &[Variable];
1183
1184    /// Accumulate gradients for momentum-based optimizers
1185    fn accumulate_gradients(&mut self, gradients: &GradientDict) -> CoreResult<()> {
1186        // Default implementation: update variable gradients
1187        for (param_name, gradient) in gradients.iter() {
1188            // Find the variable with matching name and update its gradient
1189            for var in self.variables_mut() {
1190                if var.name() == param_name {
1191                    var.set_gradient(gradient.clone())?;
1192                    break;
1193                }
1194            }
1195        }
1196        Ok(())
1197    }
1198
1199    /// Get mutable reference to variables (for default implementation)
1200    fn variables_mut(&mut self) -> &mut [Variable] {
1201        // Default implementation returns empty slice
1202        // Implementations should override this if they support accumulate_gradients
1203        &mut []
1204    }
1205}
1206
1207/// Stochastic Gradient Descent optimizer.
1208pub struct SGD {
1209    /// Variables to optimize.
1210    variables: Vec<Variable>,
1211
1212    /// Learning rate.
1213    learningrate: f64,
1214
1215    /// Momentum factor.
1216    momentum: f64,
1217
1218    /// Velocity for momentum.
1219    velocity: Vec<Option<Box<dyn ArrayProtocol>>>,
1220}
1221
1222impl SGD {
1223    /// Create a new SGD optimizer.
1224    pub fn new(learningrate: f64, momentum: Option<f64>) -> Self {
1225        Self {
1226            variables: Vec::new(),
1227            learningrate,
1228            momentum: momentum.unwrap_or(0.0),
1229            velocity: Vec::new(),
1230        }
1231    }
1232
1233    /// Set the learning rate.
1234    pub fn set_learningrate(&mut self, learningrate: f64) {
1235        self.learningrate = learningrate;
1236    }
1237}
1238
1239impl Optimizer for SGD {
1240    fn step(&mut self) -> CoreResult<()> {
1241        for (i, var) in self.variables.iter_mut().enumerate() {
1242            if let Some(grad) = var.grad_2() {
1243                let var_value = var.value();
1244
1245                // Compute update with momentum
1246                let update = if self.momentum > 0.0 {
1247                    if i >= self.velocity.len() {
1248                        self.velocity.resize_with(i + 1, || None);
1249                    }
1250
1251                    if let Some(vel) = &self.velocity[i] {
1252                        // v = momentum * v + lr * grad
1253                        let scaled_grad = multiply_by_scalar(grad.as_ref(), self.learningrate)?;
1254                        let scaled_vel = multiply_by_scalar(vel.as_ref(), self.momentum)?;
1255                        let update = add(scaled_vel.as_ref(), scaled_grad.as_ref())?;
1256                        self.velocity[i] = Some(update.clone());
1257                        update
1258                    } else {
1259                        // First iteration, just use lr * grad
1260                        let update = multiply_by_scalar(grad.as_ref(), self.learningrate)?;
1261                        self.velocity[i] = Some(update.clone());
1262                        update
1263                    }
1264                } else {
1265                    // No momentum, just use lr * grad
1266                    multiply_by_scalar(grad.as_ref(), self.learningrate)?
1267                };
1268
1269                // Update variable: var = var - update
1270                let updated_value = subtract_arrays(var_value.as_ref(), update.as_ref())?;
1271                var.set_value(updated_value);
1272            }
1273        }
1274
1275        Ok(())
1276    }
1277
1278    fn zero_grad(&mut self) {
1279        for var in &self.variables {
1280            var.tensor.node.borrow_mut().grad = None;
1281        }
1282    }
1283
1284    fn add_variable(&mut self, var: Variable) {
1285        self.variables.push(var);
1286        self.velocity.push(None);
1287    }
1288
1289    fn variables(&self) -> &[Variable] {
1290        &self.variables
1291    }
1292
1293    fn variables_mut(&mut self) -> &mut [Variable] {
1294        &mut self.variables
1295    }
1296}
1297
1298/// Adam optimizer.
1299pub struct Adam {
1300    /// Variables to optimize.
1301    variables: Vec<Variable>,
1302
1303    /// Learning rate.
1304    learningrate: f64,
1305
1306    /// Beta1 parameter (for first moment).
1307    beta1: f64,
1308
1309    /// Beta2 parameter (for second moment).
1310    beta2: f64,
1311
1312    /// Epsilon for numerical stability.
1313    epsilon: f64,
1314
1315    /// First moment estimates.
1316    m: Vec<Option<Box<dyn ArrayProtocol>>>,
1317
1318    /// Second moment estimates.
1319    v: Vec<Option<Box<dyn ArrayProtocol>>>,
1320
1321    /// Iteration counter.
1322    t: usize,
1323}
1324
1325impl Adam {
1326    /// Create a new Adam optimizer.
1327    pub fn new(
1328        learningrate: f64,
1329        beta1: Option<f64>,
1330        beta2: Option<f64>,
1331        epsilon: Option<f64>,
1332    ) -> Self {
1333        Self {
1334            variables: Vec::new(),
1335            learningrate,
1336            beta1: beta1.unwrap_or(0.9),
1337            beta2: beta2.unwrap_or(0.999),
1338            epsilon: epsilon.unwrap_or(1e-8),
1339            m: Vec::new(),
1340            v: Vec::new(),
1341            t: 0,
1342        }
1343    }
1344}
1345
1346impl Optimizer for Adam {
1347    fn step(&mut self) -> CoreResult<()> {
1348        self.t += 1;
1349
1350        for (i, var) in self.variables.iter_mut().enumerate() {
1351            if let Some(grad) = var.grad_2() {
1352                let var_value = var.value();
1353
1354                // Ensure we have space for state variables
1355                if i >= self.m.len() {
1356                    self.m.resize_with(i + 1, || None);
1357                    self.v.resize_with(i + 1, || None);
1358                }
1359
1360                // Update biased first moment estimate
1361                let m = if let Some(m_prev) = &self.m[i] {
1362                    // m = beta1 * m + (1 - beta1) * grad
1363                    let scaled_m = multiply_by_scalar(m_prev.as_ref(), self.beta1)?;
1364                    let scaled_grad = multiply_by_scalar(grad.as_ref(), 1.0 - self.beta1)?;
1365                    add(scaled_m.as_ref(), scaled_grad.as_ref())?
1366                } else {
1367                    // First iteration, just use (1 - beta1) * grad
1368                    multiply_by_scalar(grad.as_ref(), 1.0 - self.beta1)?
1369                };
1370
1371                // Update biased second moment estimate
1372                let v = if let Some(v_prev) = &self.v[i] {
1373                    // v = beta2 * v + (1 - beta2) * grad^2
1374                    let scaled_v = multiply_by_scalar(v_prev.as_ref(), self.beta2)?;
1375                    let grad_squared = multiply(grad.as_ref(), grad.as_ref())?;
1376                    let scaled_grad_sq =
1377                        multiply_by_scalar(grad_squared.as_ref(), 1.0 - self.beta2)?;
1378                    add(scaled_v.as_ref(), scaled_grad_sq.as_ref())?
1379                } else {
1380                    // First iteration, just use (1 - beta2) * grad^2
1381                    let grad_squared = multiply(grad.as_ref(), grad.as_ref())?;
1382                    multiply_by_scalar(grad_squared.as_ref(), 1.0 - self.beta2)?
1383                };
1384
1385                // Store state variables - no need to convert since we're already using Box
1386                self.m[i] = Some(m.clone());
1387                self.v[i] = Some(v.clone());
1388
1389                // Compute bias-corrected estimates
1390                let m_hat =
1391                    multiply_by_scalar(m.as_ref(), 1.0 / (1.0 - self.beta1.powi(self.t as i32)))?;
1392                let v_hat =
1393                    multiply_by_scalar(v.as_ref(), 1.0 / (1.0 - self.beta2.powi(self.t as i32)))?;
1394
1395                // Compute update: lr * m_hat / (sqrt(v_hat) + epsilon)
1396                let v_hat_sqrt = sqrt(v_hat.as_ref())?;
1397                let v_hat_sqrt_eps = add_scalar(v_hat_sqrt.as_ref(), self.epsilon)?;
1398                let update_dir = divide(m_hat.as_ref(), v_hat_sqrt_eps.as_ref())?;
1399                let update = multiply_by_scalar(update_dir.as_ref(), self.learningrate)?;
1400
1401                // Update variable: var = var - update
1402                let updated_value = subtract_arrays(var_value.as_ref(), update.as_ref())?;
1403                var.set_value(updated_value);
1404            }
1405        }
1406
1407        Ok(())
1408    }
1409
1410    fn zero_grad(&mut self) {
1411        for var in &self.variables {
1412            var.tensor.node.borrow_mut().grad = None;
1413        }
1414    }
1415
1416    fn add_variable(&mut self, var: Variable) {
1417        self.variables.push(var);
1418        self.m.push(None);
1419        self.v.push(None);
1420    }
1421
1422    fn variables(&self) -> &[Variable] {
1423        &self.variables
1424    }
1425
1426    fn variables_mut(&mut self) -> &mut [Variable] {
1427        &mut self.variables
1428    }
1429}
1430
1431// Helper functions for optimizers
1432
1433/// Extract an owned, dynamically-dimensioned copy of the array wrapped by an
1434/// `ArrayProtocol` value, regardless of whether it is concretely stored as
1435/// `Ix1`, `Ix2`, or `IxDyn`. Callers that only special-case `IxDyn` (as this
1436/// file's optimizer helpers historically did) silently fail on the very
1437/// common `Array2`/`Array1`-backed case — e.g. a `Variable` built from a 2-D
1438/// weight matrix — even though the arithmetic itself has no real dimension
1439/// restriction.
1440fn downcast_to_ixdyn<T: Clone + Send + Sync + 'static>(a: &dyn ArrayProtocol) -> Option<ArrayD<T>> {
1441    if let Some(w) = a.as_any().downcast_ref::<NdarrayWrapper<T, IxDyn>>() {
1442        return Some(w.as_array().clone());
1443    }
1444    if let Some(w) = a.as_any().downcast_ref::<NdarrayWrapper<T, Ix2>>() {
1445        return Some(w.as_array().clone().into_dyn());
1446    }
1447    if let Some(w) = a.as_any().downcast_ref::<NdarrayWrapper<T, Ix1>>() {
1448        return Some(w.as_array().clone().into_dyn());
1449    }
1450    None
1451}
1452
1453/// Multiply an array by a scalar.
1454fn multiply_by_scalar(a: &dyn ArrayProtocol, scalar: f64) -> CoreResult<Box<dyn ArrayProtocol>> {
1455    if let Some(inputarray) = downcast_to_ixdyn::<f64>(a) {
1456        let result = inputarray.mapv(|x| x * scalar);
1457        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1458    } else if let Some(inputarray) = downcast_to_ixdyn::<f32>(a) {
1459        let result = inputarray.mapv(|x| x * scalar as f32);
1460        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1461    } else if let Some(inputarray) = downcast_to_ixdyn::<i32>(a) {
1462        let result = inputarray.mapv(|x| (x as f64 * scalar) as i32);
1463        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1464    } else if let Some(inputarray) = downcast_to_ixdyn::<i64>(a) {
1465        let result = inputarray.mapv(|x| (x as f64 * scalar) as i64);
1466        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1467    } else if let Some(inputarray) = downcast_to_ixdyn::<u8>(a) {
1468        let result = inputarray.mapv(|x| (x as f64 * scalar) as u8);
1469        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1470    } else if let Some(inputarray) = downcast_to_ixdyn::<u16>(a) {
1471        let result = inputarray.mapv(|x| (x as f64 * scalar) as u16);
1472        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1473    } else if let Some(inputarray) = downcast_to_ixdyn::<u32>(a) {
1474        let result = inputarray.mapv(|x| (x as f64 * scalar) as u32);
1475        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1476    } else if let Some(inputarray) = downcast_to_ixdyn::<u64>(a) {
1477        let result = inputarray.mapv(|x| (x as f64 * scalar) as u64);
1478        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1479    } else {
1480        Err(CoreError::NotImplementedError(ErrorContext::new(
1481            "multiply_by_scalar not implemented for this array type".to_string(),
1482        )))
1483    }
1484}
1485
1486/// Subtract one array from another, returning a new array.
1487fn subtract_arrays(
1488    a: &dyn ArrayProtocol,
1489    b: &dyn ArrayProtocol,
1490) -> CoreResult<Box<dyn ArrayProtocol>> {
1491    // Perform element-wise subtraction and return a new array
1492    if let (Some(a_arr), Some(b_arr)) = (downcast_to_ixdyn::<f64>(a), downcast_to_ixdyn::<f64>(b)) {
1493        Ok(Box::new(NdarrayWrapper::new(a_arr - b_arr)) as Box<dyn ArrayProtocol>)
1494    } else if let (Some(a_arr), Some(b_arr)) =
1495        (downcast_to_ixdyn::<f32>(a), downcast_to_ixdyn::<f32>(b))
1496    {
1497        Ok(Box::new(NdarrayWrapper::new(a_arr - b_arr)) as Box<dyn ArrayProtocol>)
1498    } else if let (Some(a_arr), Some(b_arr)) =
1499        (downcast_to_ixdyn::<i32>(a), downcast_to_ixdyn::<i32>(b))
1500    {
1501        Ok(Box::new(NdarrayWrapper::new(a_arr - b_arr)) as Box<dyn ArrayProtocol>)
1502    } else if let (Some(a_arr), Some(b_arr)) =
1503        (downcast_to_ixdyn::<i64>(a), downcast_to_ixdyn::<i64>(b))
1504    {
1505        Ok(Box::new(NdarrayWrapper::new(a_arr - b_arr)) as Box<dyn ArrayProtocol>)
1506    } else if let (Some(a_arr), Some(b_arr)) =
1507        (downcast_to_ixdyn::<u8>(a), downcast_to_ixdyn::<u8>(b))
1508    {
1509        Ok(Box::new(NdarrayWrapper::new(a_arr - b_arr)) as Box<dyn ArrayProtocol>)
1510    } else if let (Some(a_arr), Some(b_arr)) =
1511        (downcast_to_ixdyn::<u16>(a), downcast_to_ixdyn::<u16>(b))
1512    {
1513        Ok(Box::new(NdarrayWrapper::new(a_arr - b_arr)) as Box<dyn ArrayProtocol>)
1514    } else if let (Some(a_arr), Some(b_arr)) =
1515        (downcast_to_ixdyn::<u32>(a), downcast_to_ixdyn::<u32>(b))
1516    {
1517        Ok(Box::new(NdarrayWrapper::new(a_arr - b_arr)) as Box<dyn ArrayProtocol>)
1518    } else if let (Some(a_arr), Some(b_arr)) =
1519        (downcast_to_ixdyn::<u64>(a), downcast_to_ixdyn::<u64>(b))
1520    {
1521        Ok(Box::new(NdarrayWrapper::new(a_arr - b_arr)) as Box<dyn ArrayProtocol>)
1522    } else {
1523        Err(CoreError::NotImplementedError(ErrorContext::new(
1524            "subtract_arrays not implemented for these array types".to_string(),
1525        )))
1526    }
1527}
1528
1529/// Element-wise square root.
1530fn sqrt(a: &dyn ArrayProtocol) -> CoreResult<Box<dyn ArrayProtocol>> {
1531    if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<f64, IxDyn>>() {
1532        let result = a_array.as_array().mapv(|x| x.sqrt());
1533        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1534    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<f32, IxDyn>>() {
1535        let result = a_array.as_array().mapv(|x| x.sqrt());
1536        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1537    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<i32, IxDyn>>() {
1538        let result = a_array.as_array().mapv(|x| (x as f64).sqrt());
1539        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1540    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<i64, IxDyn>>() {
1541        let result = a_array.as_array().mapv(|x| (x as f64).sqrt());
1542        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1543    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<u8, IxDyn>>() {
1544        let result = a_array.as_array().mapv(|x| (x as f64).sqrt());
1545        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1546    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<u16, IxDyn>>() {
1547        let result = a_array.as_array().mapv(|x| (x as f64).sqrt());
1548        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1549    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<u32, IxDyn>>() {
1550        let result = a_array.as_array().mapv(|x| (x as f64).sqrt());
1551        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1552    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<u64, IxDyn>>() {
1553        let result = a_array.as_array().mapv(|x| (x as f64).sqrt());
1554        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1555    } else {
1556        Err(CoreError::NotImplementedError(ErrorContext::new(
1557            "sqrt not implemented for this array type".to_string(),
1558        )))
1559    }
1560}
1561
1562/// Add a scalar to an array.
1563fn add_scalar(a: &dyn ArrayProtocol, scalar: f64) -> CoreResult<Box<dyn ArrayProtocol>> {
1564    if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<f64, IxDyn>>() {
1565        let result = a_array.as_array().mapv(|x| x + scalar);
1566        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1567    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<f32, IxDyn>>() {
1568        let result = a_array.as_array().mapv(|x| x + scalar as f32);
1569        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1570    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<i32, IxDyn>>() {
1571        let result = a_array.as_array().mapv(|x| x + scalar as i32);
1572        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1573    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<i64, IxDyn>>() {
1574        let result = a_array.as_array().mapv(|x| x + scalar as i64);
1575        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1576    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<u8, IxDyn>>() {
1577        let result = a_array.as_array().mapv(|x| x + scalar as u8);
1578        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1579    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<u16, IxDyn>>() {
1580        let result = a_array.as_array().mapv(|x| x + scalar as u16);
1581        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1582    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<u32, IxDyn>>() {
1583        let result = a_array.as_array().mapv(|x| x + scalar as u32);
1584        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1585    } else if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<u64, IxDyn>>() {
1586        let result = a_array.as_array().mapv(|x| x + scalar as u64);
1587        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1588    } else {
1589        Err(CoreError::NotImplementedError(ErrorContext::new(
1590            "add_scalar not implemented for this array type".to_string(),
1591        )))
1592    }
1593}
1594
1595/// Element-wise division.
1596fn divide(a: &dyn ArrayProtocol, b: &dyn ArrayProtocol) -> CoreResult<Box<dyn ArrayProtocol>> {
1597    if let (Some(a_array), Some(b_array)) = (
1598        a.as_any().downcast_ref::<NdarrayWrapper<f64, IxDyn>>(),
1599        b.as_any().downcast_ref::<NdarrayWrapper<f64, IxDyn>>(),
1600    ) {
1601        let result = a_array.as_array() / b_array.as_array();
1602        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1603    } else if let (Some(a_array), Some(b_array)) = (
1604        a.as_any().downcast_ref::<NdarrayWrapper<f32, IxDyn>>(),
1605        b.as_any().downcast_ref::<NdarrayWrapper<f32, IxDyn>>(),
1606    ) {
1607        let result = a_array.as_array() / b_array.as_array();
1608        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1609    } else if let (Some(a_array), Some(b_array)) = (
1610        a.as_any().downcast_ref::<NdarrayWrapper<i32, IxDyn>>(),
1611        b.as_any().downcast_ref::<NdarrayWrapper<i32, IxDyn>>(),
1612    ) {
1613        let result = ::ndarray::Zip::from(a_array.as_array())
1614            .and(b_array.as_array())
1615            .map_collect(|&av, &bv| av as f64 / bv as f64);
1616        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1617    } else if let (Some(a_array), Some(b_array)) = (
1618        a.as_any().downcast_ref::<NdarrayWrapper<i64, IxDyn>>(),
1619        b.as_any().downcast_ref::<NdarrayWrapper<i64, IxDyn>>(),
1620    ) {
1621        let result = ::ndarray::Zip::from(a_array.as_array())
1622            .and(b_array.as_array())
1623            .map_collect(|&av, &bv| av as f64 / bv as f64);
1624        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1625    } else if let (Some(a_array), Some(b_array)) = (
1626        a.as_any().downcast_ref::<NdarrayWrapper<u8, IxDyn>>(),
1627        b.as_any().downcast_ref::<NdarrayWrapper<u8, IxDyn>>(),
1628    ) {
1629        let result = ::ndarray::Zip::from(a_array.as_array())
1630            .and(b_array.as_array())
1631            .map_collect(|&av, &bv| av as f64 / bv as f64);
1632        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1633    } else if let (Some(a_array), Some(b_array)) = (
1634        a.as_any().downcast_ref::<NdarrayWrapper<u16, IxDyn>>(),
1635        b.as_any().downcast_ref::<NdarrayWrapper<u16, IxDyn>>(),
1636    ) {
1637        let result = ::ndarray::Zip::from(a_array.as_array())
1638            .and(b_array.as_array())
1639            .map_collect(|&av, &bv| av as f64 / bv as f64);
1640        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1641    } else if let (Some(a_array), Some(b_array)) = (
1642        a.as_any().downcast_ref::<NdarrayWrapper<u32, IxDyn>>(),
1643        b.as_any().downcast_ref::<NdarrayWrapper<u32, IxDyn>>(),
1644    ) {
1645        let result = ::ndarray::Zip::from(a_array.as_array())
1646            .and(b_array.as_array())
1647            .map_collect(|&av, &bv| av as f64 / bv as f64);
1648        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1649    } else if let (Some(a_array), Some(b_array)) = (
1650        a.as_any().downcast_ref::<NdarrayWrapper<u64, IxDyn>>(),
1651        b.as_any().downcast_ref::<NdarrayWrapper<u64, IxDyn>>(),
1652    ) {
1653        let result = ::ndarray::Zip::from(a_array.as_array())
1654            .and(b_array.as_array())
1655            .map_collect(|&av, &bv| av as f64 / bv as f64);
1656        Ok(Box::new(NdarrayWrapper::new(result)) as Box<dyn ArrayProtocol>)
1657    } else {
1658        Err(CoreError::NotImplementedError(ErrorContext::new(
1659            "divide not implemented for these array types".to_string(),
1660        )))
1661    }
1662}
1663
1664#[cfg(test)]
1665#[path = "grad_tests.rs"]
1666mod tests;