1use 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#[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 pub fn new() -> Self {
44 Self {
45 gradients: HashMap::new(),
46 }
47 }
48
49 pub fn insert(&mut self, name: String, gradient: Box<dyn ArrayProtocol>) {
51 self.gradients.insert(name, gradient);
52 }
53
54 pub fn get(&self, name: &str) -> Option<&dyn ArrayProtocol> {
56 self.gradients.get(name).map(|b| b.as_ref())
57 }
58
59 pub fn get_mut(&mut self, name: &str) -> Option<&mut Box<dyn ArrayProtocol>> {
61 self.gradients.get_mut(name)
62 }
63
64 pub fn iter(&self) -> impl Iterator<Item = (&String, &Box<dyn ArrayProtocol>)> {
66 self.gradients.iter()
67 }
68
69 pub fn merge(&mut self, other: GradientDict) {
71 for (name, gradient) in other.gradients {
72 self.gradients.insert(name, gradient);
73 }
74 }
75
76 pub fn is_empty(&self) -> bool {
78 self.gradients.is_empty()
79 }
80
81 pub fn len(&self) -> usize {
83 self.gradients.len()
84 }
85
86 pub fn clear(&mut self) {
88 self.gradients.clear();
89 }
90
91 pub fn keys(&self) -> impl Iterator<Item = &String> {
93 self.gradients.keys()
94 }
95
96 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#[allow(dead_code)]
110fn boxed_to_rc(boxed: Box<dyn ArrayProtocol>) -> Rc<dyn ArrayProtocol> {
111 Rc::from(boxed)
119}
120
121#[allow(dead_code)]
123fn box_to_rc_array_protocol(boxed: Box<dyn ArrayProtocol>) -> Rc<dyn ArrayProtocol> {
124 boxed_to_rc(boxed)
125}
126
127#[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 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 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 if let Some(a_array) = a.as_any().downcast_ref::<NdarrayWrapper<f64, IxDyn>>() {
175 let array = a_array.as_array();
176 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 Ok(Box::new(NdarrayWrapper::new(array.clone())) as Box<dyn ArrayProtocol>)
184 } else {
185 let inputshape = array.shape();
187 let _ndim_diff = shape.len().saturating_sub(inputshape.len());
188
189 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 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 let ones = ArrayD::<f64>::ones(IxDyn(shape));
235 Ok(Box::new(NdarrayWrapper::new(ones)) as Box<dyn ArrayProtocol>)
236 }
237}
238
239#[derive(Clone)]
241struct Node {
242 value: Rc<dyn ArrayProtocol>,
244
245 grad: Option<Rc<dyn ArrayProtocol>>,
247
248 op: Option<String>,
250
251 inputs: Vec<GradientTensor>,
253
254 requiresgrad: bool,
256
257 is_leaf: bool,
259}
260
261impl Node {
262 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 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#[derive(Clone)]
293pub struct GradientTensor {
294 node: Rc<RefCell<Node>>,
296}
297
298impl GradientTensor {
299 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 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 pub fn value(&self) -> Rc<dyn ArrayProtocol> {
320 self.node.borrow().value.clone()
321 }
322
323 pub fn grad_2(&self) -> Option<Rc<dyn ArrayProtocol>> {
325 self.node.borrow().grad.clone()
326 }
327
328 pub fn requiresgrad(&self) -> bool {
330 self.node.borrow().requiresgrad
331 }
332
333 pub fn set_requiresgrad(&mut self, requiresgrad: bool) {
335 self.node.borrow_mut().requiresgrad = requiresgrad;
336 }
337
338 pub fn is_leaf(&self) -> bool {
340 self.node.borrow().is_leaf
341 }
342
343 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 pub fn set_value(&mut self, newvalue: Rc<dyn ArrayProtocol>) {
351 self.node.borrow_mut().grad = None; self.node.borrow_mut().value = newvalue;
353 }
354
355 pub fn backward(&self) -> CoreResult<()> {
357 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 self.backward_with_grad(grad)
371 }
372
373 fn backward_with_grad(&self, grad: Rc<dyn ArrayProtocol>) -> CoreResult<()> {
375 self.node.borrow_mut().grad = Some(grad.clone());
377
378 let mut visited = HashSet::new();
380 let mut topo = Vec::new();
381
382 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 for input in &tensor.node.borrow().inputs {
394 build_topo(input, visited, topo);
395 }
396
397 topo.push(tensor.clone());
399 }
400 }
401
402 build_topo(self, &mut visited, &mut topo);
404
405 for node in topo.iter().rev() {
407 if !node.requiresgrad() {
409 continue;
410 }
411
412 let node_grad = match node.grad_2() {
414 Some(g) => g,
415 None => continue, };
417
418 if node.is_leaf() {
420 continue;
421 }
422
423 let op = match &node.node.borrow().op {
425 Some(op) => op.clone(),
426 None => continue, };
428
429 let inputs = node.node.borrow().inputs.clone();
430
431 match op.as_str() {
433 "add" => {
434 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 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 if inputs.len() == 2 => {
452 let (a, b) = (&inputs[0], &inputs[1]);
453
454 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 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 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 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 if inputs.len() == 2 => {
491 let (a, b) = (&inputs[0], &inputs[1]);
492
493 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 let b_t = b_array_val.t();
508
509 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 let grad_a_val = grad_out_2d.dot(&b_t_2d);
537
538 let grad_a_dyn = grad_a_val.into_dyn();
540 let grad_a = NdarrayWrapper::new(grad_a_dyn);
541
542 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 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 a_node.grad = Some(Rc::new(grad_a));
560 }
561 }
562 }
563
564 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 let a_t = a_array_val.t();
579
580 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 let grad_b_val = a_t_2d.dot(&grad_out_2d);
608
609 let grad_b_dyn = grad_b_val.into_dyn();
611 let grad_b = NdarrayWrapper::new(grad_b_dyn);
612
613 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 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 b_node.grad = Some(Rc::new(grad_b));
631 }
632 }
633 }
634 }
635 "subtract"
636 if inputs.len() == 2 => {
638 let (a, b) = (&inputs[0], &inputs[1]);
639
640 if a.requiresgrad() {
642 let mut a_node = a.node.borrow_mut();
643 if let Some(a_grad) = &a_node.grad {
644 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 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 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 if inputs.len() == 2 => {
671 let (a, b) = (&inputs[0], &inputs[1]);
672
673 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 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 if b.requiresgrad() {
691 let a_value = a.value();
692 let b_value = b.value();
693
694 if let Ok(b_squared) = multiply(b_value.as_ref(), b_value.as_ref()) {
696 if let Ok(grad_times_a) =
698 multiply(node_grad.as_ref(), a_value.as_ref())
699 {
700 if let Ok(div_result) =
702 divide(grad_times_a.as_ref(), b_squared.as_ref())
703 {
704 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 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 if inputs.len() == 1 => {
730 let input = &inputs[0];
731
732 if input.requiresgrad() {
733 let sigmoid_value = node.value();
735
736 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 if let Ok(sigmoid_deriv) =
743 multiply(sigmoid_value.as_ref(), one_minus_sigmoid.as_ref())
744 {
745 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 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 if inputs.len() == 1 => {
771 let input = &inputs[0];
772
773 if input.requiresgrad() {
774 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 if let Ok(grad_input) =
784 multiply_by_scalar(node_grad.as_ref(), 1.0 / n_elements)
785 {
786 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 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 }
812 }
813 }
814
815 Ok(())
816 }
817
818 pub fn detach(&self) -> Self {
820 GradientTensor::new(self.value(), false)
821 }
822}
823
824#[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 let result = add(a_value.as_ref(), b_value.as_ref())?;
833
834 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#[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 let result = multiply(a_value.as_ref(), b_value.as_ref())?;
851
852 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#[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 let result = matmul(a_value.as_ref(), b_value.as_ref())?;
869
870 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#[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 let result = subtract(a_value.as_ref(), b_value.as_ref())?;
887
888 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#[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 let result = divide(a_value.as_ref(), b_value.as_ref())?;
905
906 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#[allow(dead_code)]
917pub fn grad_sigmoid(a: &GradientTensor) -> CoreResult<GradientTensor> {
918 let a_value = a.value();
919
920 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#[allow(dead_code)]
956pub fn grad_mean(a: &GradientTensor) -> CoreResult<GradientTensor> {
957 let a_value = a.value();
958
959 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
1100pub struct Variable {
1102 tensor: GradientTensor,
1104
1105 name: String,
1107}
1108
1109impl Variable {
1110 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 pub const fn tensor(&self) -> &GradientTensor {
1125 &self.tensor
1126 }
1127
1128 pub fn value(&self) -> Rc<dyn ArrayProtocol> {
1130 self.tensor.value()
1131 }
1132
1133 pub fn grad_2(&self) -> Option<Rc<dyn ArrayProtocol>> {
1135 self.tensor.grad_2()
1136 }
1137
1138 pub fn name(&self) -> &str {
1140 &self.name
1141 }
1142
1143 pub fn set_gradient(&mut self, gradient: Box<dyn ArrayProtocol>) -> CoreResult<()> {
1145 let gradient_rc = self.box_to_rc(gradient);
1147
1148 self.tensor.node.borrow_mut().grad = Some(gradient_rc);
1150 Ok(())
1151 }
1152
1153 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 fn box_to_rc(&self, boxed: Box<dyn ArrayProtocol>) -> Rc<dyn ArrayProtocol> {
1161 Rc::from(boxed)
1167 }
1168}
1169
1170pub trait Optimizer {
1172 fn step(&mut self) -> CoreResult<()>;
1174
1175 fn zero_grad(&mut self);
1177
1178 fn add_variable(&mut self, var: Variable);
1180
1181 fn variables(&self) -> &[Variable];
1183
1184 fn accumulate_gradients(&mut self, gradients: &GradientDict) -> CoreResult<()> {
1186 for (param_name, gradient) in gradients.iter() {
1188 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 fn variables_mut(&mut self) -> &mut [Variable] {
1201 &mut []
1204 }
1205}
1206
1207pub struct SGD {
1209 variables: Vec<Variable>,
1211
1212 learningrate: f64,
1214
1215 momentum: f64,
1217
1218 velocity: Vec<Option<Box<dyn ArrayProtocol>>>,
1220}
1221
1222impl SGD {
1223 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 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 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 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 let update = multiply_by_scalar(grad.as_ref(), self.learningrate)?;
1261 self.velocity[i] = Some(update.clone());
1262 update
1263 }
1264 } else {
1265 multiply_by_scalar(grad.as_ref(), self.learningrate)?
1267 };
1268
1269 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
1298pub struct Adam {
1300 variables: Vec<Variable>,
1302
1303 learningrate: f64,
1305
1306 beta1: f64,
1308
1309 beta2: f64,
1311
1312 epsilon: f64,
1314
1315 m: Vec<Option<Box<dyn ArrayProtocol>>>,
1317
1318 v: Vec<Option<Box<dyn ArrayProtocol>>>,
1320
1321 t: usize,
1323}
1324
1325impl Adam {
1326 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 if i >= self.m.len() {
1356 self.m.resize_with(i + 1, || None);
1357 self.v.resize_with(i + 1, || None);
1358 }
1359
1360 let m = if let Some(m_prev) = &self.m[i] {
1362 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 multiply_by_scalar(grad.as_ref(), 1.0 - self.beta1)?
1369 };
1370
1371 let v = if let Some(v_prev) = &self.v[i] {
1373 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 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 self.m[i] = Some(m.clone());
1387 self.v[i] = Some(v.clone());
1388
1389 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 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 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
1431fn 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
1453fn 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
1486fn subtract_arrays(
1488 a: &dyn ArrayProtocol,
1489 b: &dyn ArrayProtocol,
1490) -> CoreResult<Box<dyn ArrayProtocol>> {
1491 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
1529fn 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
1562fn 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
1595fn 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;