use crate::op::binary;
use crate::tape::sealed::ThisThreadTape;
use crate::IndexT;
use crate::ad::AD;
use crate::op::info::OpInfo;
use crate::op::id::{
MUL_CV_OP,
MUL_VC_OP,
MUL_VV_OP,
};
#[cfg(doc)]
use crate::op::info::{
ForwardOne,
ReverseOne,
};
binary::binary_rust_src!(Mul, *);
binary::eval_binary_forward_0!(Mul, *);
fn mul_cv_forward_1 <V, E>(
_var_zero : &Vec<E> ,
var_one : &mut Vec<E> ,
con : &Vec<V> ,
_flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
for<'a> &'a V : std::ops::Mul<&'a E, Output = E> ,
{
debug_assert!( arg.len() == 2);
let lhs = arg[0] as usize;
let rhs = arg[1] as usize;
var_one[ res ] = &con[lhs] * &var_one[rhs];
}
fn mul_vc_forward_1 <V, E>(
_var_zero : &Vec<E> ,
var_one : &mut Vec<E> ,
con : &Vec<V> ,
_flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
for<'a> &'a E : std::ops::Mul<&'a V, Output = E> ,
{
debug_assert!( arg.len() == 2);
let lhs = arg[0] as usize;
let rhs = arg[1] as usize;
var_one[ res ] = &var_one[lhs] * &con[rhs];
}
fn mul_vv_forward_1 <V, E>(
var_zero : &Vec<E> ,
var_one : &mut Vec<E> ,
_con : &Vec<V> ,
_flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
for<'a> &'a E : std::ops::Add<&'a E, Output = E> ,
for<'a> &'a E : std::ops::Mul<&'a E, Output = E> ,
{
debug_assert!( arg.len() == 2);
let lhs = arg[0] as usize;
let rhs = arg[1] as usize;
let term1 = &var_zero[lhs] * &var_one[rhs];
let term2 = &var_one[lhs] * &var_zero[rhs];
var_one[res] = &term1 + &term2;
}
fn mul_cv_reverse_1 <V, E>(
_var_zero : &Vec<E> ,
var_one : &mut Vec<E> ,
con : &Vec<V> ,
_flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
for<'a> E : std::ops::AddAssign<&'a E> ,
for<'a> &'a E : std::ops::Mul<&'a V, Output = E> ,
{
debug_assert!( arg.len() == 2);
let lhs = arg[0] as usize;
let rhs = arg[1] as usize;
let term = &var_one[res] * &con[lhs];
var_one[rhs] += &term;
}
fn mul_vc_reverse_1 <V, E>(
_var_zero : &Vec<E> ,
var_one : &mut Vec<E> ,
con : &Vec<V> ,
_flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
for<'a> E : std::ops::AddAssign<&'a E> ,
for<'a> &'a E : std::ops::Mul<&'a V, Output = E> ,
{
debug_assert!( arg.len() == 2);
let lhs = arg[0] as usize;
let rhs = arg[1] as usize;
let term = &var_one[res] * &con[rhs];
var_one[lhs] += &term;
}
fn mul_vv_reverse_1 <V, E>(
var_zero : &Vec<E> ,
var_one : &mut Vec<E> ,
_con : &Vec<V> ,
_flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
for<'a> E : std::ops::AddAssign<&'a E> ,
for<'a> &'a E : std::ops::Mul<&'a E, Output = E> ,
{
debug_assert!( arg.len() == 2);
let lhs = arg[0] as usize;
let rhs = arg[1] as usize;
let term = &var_one[res] * &var_zero[rhs];
var_one[lhs] += &term;
let term = &var_one[res] * &var_zero[lhs];
var_one[rhs] += &term;
}
pub fn set_op_info<V>( op_info_vec : &mut Vec< OpInfo<V> > )
where
for<'a> V : std::ops::AddAssign<&'a V> ,
for<'a> &'a V : std::ops::Add<&'a AD<V>, Output = AD<V> > ,
for<'a> &'a V : std::ops::Add<&'a V, Output = V> ,
for<'a> &'a V : std::ops::Mul<&'a AD<V>, Output = AD<V> > ,
for<'a> &'a V : std::ops::Mul<&'a V, Output = V> ,
V : Clone + ThisThreadTape ,
{
op_info_vec[MUL_CV_OP as usize] = OpInfo{
name : "mul_cv",
forward_0_value : mul_cv_forward_0::<V, V>,
forward_0_ad : mul_cv_forward_0::<V, AD<V> >,
forward_1_value : mul_cv_forward_1::<V, V>,
forward_1_ad : mul_cv_forward_1::<V, AD<V> >,
reverse_1_value : mul_cv_reverse_1::<V, V>,
reverse_1_ad : mul_cv_reverse_1::<V, AD<V> >,
arg_var_index : binary::binary_cv_arg_var_index,
rust_src : mul_cv_rust_src,
};
op_info_vec[MUL_VC_OP as usize] = OpInfo{
name : "mul_vc",
forward_0_value : mul_vc_forward_0::<V, V>,
forward_0_ad : mul_vc_forward_0::<V, AD<V> >,
forward_1_value : mul_vc_forward_1::<V, V>,
forward_1_ad : mul_vc_forward_1::<V, AD<V> >,
reverse_1_value : mul_vc_reverse_1::<V, V>,
reverse_1_ad : mul_vc_reverse_1::<V, AD<V> >,
arg_var_index : binary::binary_vc_arg_var_index,
rust_src : mul_vc_rust_src,
};
op_info_vec[MUL_VV_OP as usize] = OpInfo{
name : "mul_vv",
forward_0_value : mul_vv_forward_0::<V, V>,
forward_0_ad : mul_vv_forward_0::<V, AD<V> >,
forward_1_value : mul_vv_forward_1::<V, V>,
forward_1_ad : mul_vv_forward_1::<V, AD<V> >,
reverse_1_value : mul_vv_reverse_1::<V, V>,
reverse_1_ad : mul_vv_reverse_1::<V, AD<V> >,
arg_var_index : binary::binary_vv_arg_var_index,
rust_src : mul_vv_rust_src,
};
}