use std::sync::RwLock;
use crate::op::info::{
OpInfo,
panic_rust_src,
};
use crate::atom::{
AtomForwardZeroValue,
AtomForwardZeroAD,
AtomForwardOneValue,
AtomForwardOneAD,
AtomReverseOneValue,
AtomReverseOneAD,
sealed::AtomEvalVec,
};
use crate::op::id::{
CALL_OP,
CALL_RES_OP,
};
use crate::{
AD,
IndexT,
AtomEval,
ThisThreadTapePublic,
ad_from_value,
};
fn extract_call_arg<'a>(
flag : &'a Vec<bool> ,
arg : &'a [IndexT] ,
) -> (
usize , // atom_id
IndexT , // call_info
usize , // call_n_arg
usize , // call_n_res
bool , // trace
&'a [bool] , // is_arg_var
&'a [bool] , // is_res_var
) {
let atom_id = arg[0] as usize;
let call_info = arg[1];
let call_n_arg = arg[2] as usize;
let call_n_res = arg[3] as usize;
let trace = flag[ arg[4] as usize ];
let mut begin = (arg[4] as usize) + 1;
let mut end = begin + call_n_arg;
let is_arg_var = &flag[begin .. end];
begin = end;
end = begin + call_n_res;
let is_res_var = &flag[begin .. end];
(
atom_id,
call_info,
call_n_arg,
call_n_res,
trace,
is_arg_var,
is_res_var,
)
}
fn call_domain_zero_value<'a, 'b, V>(
var_zero : &'a Vec<V> ,
con : &'a Vec<V> ,
arg : &'b [IndexT] ,
call_n_arg : usize ,
is_arg_var : &'b [bool] ,
) -> Vec<&'a V>
{
let mut call_domain_zero : Vec<&V> = Vec::with_capacity( call_n_arg );
for i_arg in 0 .. call_n_arg {
let index = arg[i_arg + 5] as usize;
if is_arg_var[i_arg] {
call_domain_zero.push( &var_zero[index] );
} else {
call_domain_zero.push( &con[index] );
}
}
call_domain_zero
}
fn call_domain_acon<'a, 'b, V>(
con : &'a Vec<V> ,
arg : &'b [IndexT] ,
call_n_arg : usize ,
is_arg_var : &'b [bool] ,
) -> Vec< AD<V> >
where
V : Clone,
{
let mut acon : Vec< AD<V> > = Vec::new();
for i_arg in 0 .. call_n_arg {
if ! is_arg_var[i_arg] {
let index = arg[i_arg + 5] as usize;
acon.push( ad_from_value( con[index].clone() ) );
}
}
acon
}
fn call_domain_zero_ad<'a, 'b, V>(
avar_zero : &'a Vec< AD<V> > ,
acon : &'a Vec< AD<V> > ,
arg : &'b [IndexT] ,
call_n_arg : usize ,
is_arg_var : &'b [bool] ,
) -> Vec<&'a AD<V> >
{
let mut call_domain_zero : Vec<& AD<V> > = Vec::with_capacity( call_n_arg );
let mut i_con : usize = 0;
for i_arg in 0 .. call_n_arg {
if is_arg_var[i_arg] {
let index = arg[i_arg + 5] as usize;
call_domain_zero.push( &avar_zero[index] );
} else {
call_domain_zero.push( &acon[i_con] );
i_con += 1;
}
}
call_domain_zero
}
fn call_forward_0_value<V> (
var_zero : &mut Vec<V> ,
con : &Vec<V> ,
flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
V : AtomEvalVec,
{ let (
atom_id,
call_info,
call_n_arg,
call_n_res,
trace,
is_arg_var,
is_res_var,
) = extract_call_arg(flag, arg);
let call_domain_zero = call_domain_zero_value(
var_zero, con, arg, call_n_arg, is_arg_var
);
let forward_zero_value : AtomForwardZeroValue<V>;
{ let rw_lock : &RwLock< Vec< AtomEval<V> > > = AtomEvalVec::get();
let read_lock = rw_lock.read();
assert!( read_lock.is_ok() );
let atom_eval_vec = read_lock.unwrap();
let atom_eval = &atom_eval_vec[atom_id];
forward_zero_value = atom_eval.forward_zero_value.clone();
}
let mut call_range_zero = forward_zero_value(
&call_domain_zero, call_info, trace
);
let mut j_res = 0;
call_range_zero.reverse();
for i_res in (0 .. call_n_res).rev() {
let range_i = call_range_zero.pop();
debug_assert!( range_i.is_some() );
if is_res_var[i_res] {
var_zero[res + j_res] = range_i.unwrap();
j_res += 1;
}
}
}
fn call_forward_1_value<V> (
var_zero : &Vec<V> ,
var_one : &mut Vec<V> ,
con : &Vec<V> ,
flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
V : AtomEvalVec + From<f32>,
{ let (
atom_id,
call_info,
call_n_arg,
call_n_res,
trace,
is_arg_var,
is_res_var,
) = extract_call_arg(flag, arg);
let call_domain_zero = call_domain_zero_value(
var_zero, con, arg, call_n_arg, is_arg_var
);
let name : &'static str;
let forward_zero_value : AtomForwardZeroValue<V>;
let forward_one_value : Option< AtomForwardOneValue<V> >;
{ let rw_lock : &RwLock< Vec< AtomEval<V> > > = AtomEvalVec::get();
let read_lock = rw_lock.read();
assert!( read_lock.is_ok() );
let atom_eval_vec = read_lock.unwrap();
let atom_eval = &atom_eval_vec[atom_id];
forward_zero_value = atom_eval.forward_zero_value.clone();
name = atom_eval.name;
forward_one_value = atom_eval.forward_one_value.clone();
}
if forward_one_value.is_none() {
panic!(
"{} : forward_one_value not implemented for this atomic function",
name,
);
}
let forward_one_value = forward_one_value.unwrap();
forward_zero_value(&call_domain_zero, call_info, trace);
let zero_v : V = 0f32.into();
let mut call_domain_one : Vec<&V> = Vec::with_capacity( call_n_arg );
for i_arg in 0 .. call_n_arg {
let index = arg[i_arg + 5] as usize;
if is_arg_var[i_arg] {
call_domain_one.push( &var_one[index] );
} else {
call_domain_one.push( &zero_v );
}
}
let mut call_range_one = forward_one_value(
&call_domain_zero, call_domain_one, call_info, trace
);
let mut j_res = 0;
call_range_one.reverse();
for i_res in (0 .. call_n_res).rev() {
let range_i = call_range_one.pop();
debug_assert!( range_i.is_some() );
if is_res_var[i_res] {
var_one[res + j_res] = range_i.unwrap();
j_res += 1;
}
}
}
fn call_reverse_1_value<V> (
var_zero : &Vec<V> ,
var_one : &mut Vec<V> ,
con : &Vec<V> ,
flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
for<'a> V : AtomEvalVec + std::ops::AddAssign<&'a V> + From<f32>,
{ let (
atom_id,
call_info,
call_n_arg,
call_n_res,
trace,
is_arg_var,
is_res_var,
) = extract_call_arg(flag, arg);
let call_domain_zero = call_domain_zero_value(
var_zero, con, arg, call_n_arg, is_arg_var
);
let name : &'static str;
let reverse_one_value : Option< AtomReverseOneValue<V> >;
{ let rw_lock : &RwLock< Vec< AtomEval<V> > > = AtomEvalVec::get();
let read_lock = rw_lock.read();
assert!( read_lock.is_ok() );
let atom_eval_vec = read_lock.unwrap();
let atom_eval = &atom_eval_vec[atom_id];
name = atom_eval.name;
reverse_one_value = atom_eval.reverse_one_value.clone();
}
if reverse_one_value.is_none() {
panic!(
"{}: reverse_one_value not implemented for this atomic function",
name,
);
}
let reverse_one_value = reverse_one_value.unwrap();
let zero_v : V = 0f32.into();
let mut call_range_one : Vec<&V> = Vec::with_capacity( call_n_res );
let mut j_res = 0;
for i_res in 0 .. call_n_res {
if is_res_var[i_res] {
call_range_one.push( &var_one[res + j_res] );
j_res += 1;
} else {
call_range_one.push( &zero_v );
}
}
let call_domain_one = reverse_one_value(
&call_domain_zero, call_range_one, call_info, trace
);
for i_arg in 0 .. call_n_arg {
let index = arg[i_arg + 5] as usize;
if is_arg_var[i_arg] {
var_one[index] += &call_domain_one[i_arg];
}
}
}
fn call_forward_0_ad<V> (
avar_zero : &mut Vec< AD<V> > ,
con : &Vec<V> ,
flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
V : Clone + AtomEvalVec,
{ let (
atom_id,
call_info,
call_n_arg,
call_n_res,
trace,
is_arg_var,
is_res_var,
) = extract_call_arg(flag, arg);
let acon = call_domain_acon(con, arg, call_n_arg, is_arg_var);
let call_adomain_zero = call_domain_zero_ad(
avar_zero, &acon, arg, call_n_arg, is_arg_var
);
let name : &'static str;
let forward_zero_ad : Option< AtomForwardZeroAD<V> >;
{ let rw_lock : &RwLock< Vec< AtomEval<V> > > = AtomEvalVec::get();
let read_lock = rw_lock.read();
assert!( read_lock.is_ok() );
let atom_eval_vec = read_lock.unwrap();
let atom_eval = &atom_eval_vec[atom_id];
name = atom_eval.name;
forward_zero_ad = atom_eval.forward_zero_ad.clone();
}
if forward_zero_ad.is_none() {
panic!(
"{} : forward_zero_ad is not implemented for this atomic function",
name,
);
}
let forward_zero_ad = forward_zero_ad.unwrap();
let mut call_arange_zero = forward_zero_ad(
&call_adomain_zero, call_info, trace
);
let mut j_res = 0;
call_arange_zero.reverse();
for i_res in (0 .. call_n_res).rev() {
let arange_i = call_arange_zero.pop();
debug_assert!( arange_i.is_some() );
if is_res_var[i_res] {
avar_zero[res + j_res] = arange_i.unwrap();
j_res += 1;
}
}
}
fn call_forward_1_ad<V> (
avar_zero : &Vec< AD<V> > ,
avar_one : &mut Vec< AD<V> > ,
con : &Vec<V> ,
flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
V : From<f32> + Clone + AtomEvalVec ,
{ let (
atom_id,
call_info,
call_n_arg,
call_n_res,
trace,
is_arg_var,
is_res_var,
) = extract_call_arg(flag, arg);
let acon = call_domain_acon(con, arg, call_n_arg, is_arg_var);
let call_adomain_zero = call_domain_zero_ad(
avar_zero, &acon, arg, call_n_arg, is_arg_var
);
let name : &'static str;
let forward_zero_ad : Option< AtomForwardZeroAD<V> >;
let forward_one_ad : Option< AtomForwardOneAD<V> >;
{ let rw_lock : &RwLock< Vec< AtomEval<V> > > = AtomEvalVec::get();
let read_lock = rw_lock.read();
assert!( read_lock.is_ok() );
let atom_eval_vec = read_lock.unwrap();
let atom_eval = &atom_eval_vec[atom_id];
name = atom_eval.name;
forward_zero_ad = atom_eval.forward_zero_ad.clone();
forward_one_ad = atom_eval.forward_one_ad.clone();
}
if forward_zero_ad.is_none() {
panic!(
"{} : forward_zero_ad is not implemented for this atomic function",
name,
);
}
let forward_zero_ad = forward_zero_ad.unwrap();
if forward_one_ad.is_none() {
panic!(
"{} : forward_one_ad is not implemented for this atomic function",
name,
);
}
let forward_one_ad = forward_one_ad.unwrap();
forward_zero_ad(&call_adomain_zero, call_info, trace);
let zero_v : V = 0.0f32.into();
let azero = ad_from_value(zero_v);
let mut call_adomain_one : Vec<& AD<V> > = Vec::with_capacity(call_n_arg);
for i_arg in 0 .. call_n_arg {
let index = arg[i_arg + 5] as usize;
if is_arg_var[i_arg] {
call_adomain_one.push( &avar_one[index] );
} else {
call_adomain_one.push( &azero );
}
}
let mut call_arange_one = forward_one_ad(
&call_adomain_zero, call_adomain_one, call_info, trace
);
let mut j_res = 0;
call_arange_one.reverse();
for i_res in (0 .. call_n_res).rev() {
let arange_i = call_arange_one.pop();
debug_assert!( arange_i.is_some() );
if is_res_var[i_res] {
avar_one[res + j_res] = arange_i.unwrap();
j_res += 1;
}
}
}
fn call_reverse_1_ad<V> (
avar_zero : &Vec< AD<V> > ,
avar_one : &mut Vec< AD<V> > ,
con : &Vec<V> ,
flag : &Vec<bool> ,
arg : &[IndexT] ,
res : usize )
where
V : AtomEvalVec + Clone + From<f32>,
for<'a> AD<V> : std::ops::AddAssign<&'a AD<V> >,
{ let (
atom_id,
call_info,
call_n_arg,
call_n_res,
trace,
is_arg_var,
is_res_var,
) = extract_call_arg(flag, arg);
let acon = call_domain_acon(con, arg, call_n_arg, is_arg_var);
let call_adomain_zero = call_domain_zero_ad(
avar_zero, &acon, arg, call_n_arg, is_arg_var
);
let name : &'static str;
let reverse_one_ad : Option< AtomReverseOneAD<V> >;
{ let rw_lock : &RwLock< Vec< AtomEval<V> > > = AtomEvalVec::get();
let read_lock = rw_lock.read();
assert!( read_lock.is_ok() );
let atom_eval_vec = read_lock.unwrap();
let atom_eval = &atom_eval_vec[atom_id];
name = atom_eval.name;
reverse_one_ad = atom_eval.reverse_one_ad.clone();
}
if reverse_one_ad.is_none() {
panic!(
"{}: reverse_one_ad not implemented for this atomic function",
name,
);
}
let reverse_one_ad = reverse_one_ad.unwrap();
let zero_v : V = 0f32.into();
let azero = ad_from_value(zero_v);
let mut call_arange_one : Vec<& AD<V>> = Vec::with_capacity( call_n_res );
let mut j_res = 0;
for i_res in 0 .. call_n_res {
if is_res_var[i_res] {
call_arange_one.push( &avar_one[res + j_res] );
j_res += 1;
} else {
call_arange_one.push( &azero );
}
}
let call_adomain_one = reverse_one_ad(
&call_adomain_zero, call_arange_one, call_info, trace
);
for i_arg in 0 .. call_n_arg {
let index = arg[i_arg + 5] as usize;
if is_arg_var[i_arg] {
avar_one[index] += &call_adomain_one[i_arg];
}
}
}
fn call_arg_var_index(
arg_var_index : &mut Vec<IndexT>,
flag : &Vec<bool>,
arg : &[IndexT]
)
{
let call_n_arg = arg[2] as usize;
let begin = arg[3] as usize;
let end = begin + call_n_arg;
let is_var = &flag[begin .. end];
let zero_t = 0 as IndexT;
arg_var_index.resize(0, zero_t);
for call_i_arg in 0 .. call_n_arg {
if is_var[call_i_arg] {
arg_var_index.push( arg[5 + call_i_arg] );
}
}
assert_ne!( arg_var_index.len() , 0 );
}
pub(crate) fn set_op_info<V>( op_info_vec : &mut Vec< OpInfo<V> > )
where
V : Clone + From<f32> + AtomEvalVec + ThisThreadTapePublic,
for<'a> V : std::ops::AddAssign<&'a V> ,
{
op_info_vec[CALL_OP as usize] = OpInfo{
name : "call" ,
forward_0_value : call_forward_0_value::<V>,
forward_0_ad : call_forward_0_ad::<V>,
forward_1_value : call_forward_1_value::<V>,
forward_1_ad : call_forward_1_ad::<V>,
reverse_1_value : call_reverse_1_value::<V>,
reverse_1_ad : call_reverse_1_ad::<V>,
arg_var_index : call_arg_var_index,
rust_src : panic_rust_src,
};
op_info_vec[CALL_RES_OP as usize] = OpInfo{
name : "call_res" ,
forward_0_value : no_op_zero::<V, V>,
forward_0_ad : no_op_zero::<V, AD<V> >,
forward_1_value : no_op_one::<V, V>,
forward_1_ad : no_op_one::<V, AD<V> >,
reverse_1_value : no_op_one::<V, V>,
reverse_1_ad : no_op_one::<V, AD<V> >,
arg_var_index : no_op_arg_var_index,
rust_src : panic_rust_src,
};
}
fn no_op_zero<V, E>(
_var_zero : &mut Vec<E> ,
_con : &Vec<V> ,
_flag : &Vec<bool> ,
_arg : &[IndexT] ,
_res : usize ,
) { }
fn no_op_one<V, E>(
_var_zero : &Vec<E> ,
_var_one : &mut Vec<E> ,
_con : &Vec<V> ,
_flag : &Vec<bool> ,
_arg : &[IndexT] ,
_res : usize ,
) { }
fn no_op_arg_var_index(
arg_var_index : &mut Vec<IndexT> ,
_flag : &Vec<bool> ,
_arg : &[IndexT] ,
) {
let zero_t = 0 as IndexT;
arg_var_index.resize(0, zero_t);
}