use std::sync::RwLock;
use std::thread::LocalKey;
use std::cell::RefCell;
use crate::op::id::CALL_OP;
use crate::op::id::CALL_RES_OP;
use crate::tape::Tape;
use crate::tape::sealed::ThisThreadTape;
use crate::{
IndexT,
AD,
ad_from_vector,
AtomEvalVecPublic,
ThisThreadTapePublic,
};
#[cfg(doc)]
use crate::{
doc_generic_v,
ADfn,
};
#[cfg(doc)]
use crate::adfn::{
forward_zero::doc_forward_zero,
forward_one::doc_forward_one,
reverse_one::doc_reverse_one,
};
pub type AtomForwardZeroValue<V> = fn(
_domain_zero : &Vec<&V> ,
_call_info : IndexT ,
_trace : bool ,
) -> Vec<V> ;
pub type AtomForwardOneValue<V> = fn(
_domain_zero : &Vec<&V> ,
_domain_one : Vec<&V> ,
_call_info : IndexT ,
_trace : bool ,
) -> Vec<V> ;
pub type AtomReverseOneValue<V> = fn(
_domain_zero : &Vec<&V> ,
_range_one : Vec<&V> ,
_call_info : IndexT ,
_trace : bool ,
) -> Vec<V> ;
pub type AtomForwardDepend = fn(
_is_var_domain : &Vec<bool> ,
_call_info : IndexT ,
_trace : bool ,
)-> Vec<bool>;
pub type AtomForwardZeroAD<V> = fn(
_domain_zero : &Vec<& AD<V> > ,
_call_info : IndexT ,
_trace : bool ,
) -> Vec< AD<V> > ;
pub type AtomForwardOneAD<V> = fn(
_domain_zero : &Vec<& AD<V> > ,
_domain_one : Vec<& AD<V> > ,
_call_info : IndexT ,
_trace : bool ,
) -> Vec< AD<V> > ;
pub type AtomReverseOneAD<V> = fn(
_domain_zero : &Vec<& AD<V> > ,
_range_one : Vec<& AD<V> > ,
_call_info : IndexT ,
_trace : bool ,
) -> Vec< AD<V> > ;
pub struct AtomEval<V> {
pub name : &'static str ,
pub forward_depend : AtomForwardDepend ,
pub forward_zero_value : AtomForwardZeroValue::<V> ,
pub forward_zero_ad : Option< AtomForwardZeroAD::<V> >,
pub forward_one_value : Option< AtomForwardOneValue::<V> > ,
pub forward_one_ad : Option< AtomForwardOneAD::<V> > ,
pub reverse_one_value : Option< AtomReverseOneValue::<V> > ,
pub reverse_one_ad : Option< AtomReverseOneAD::<V> > ,
}
pub (crate) mod sealed {
use std::sync::RwLock;
use super::AtomEval;
pub trait AtomEvalVec
where
Self : Sized + 'static,
{ fn get() -> &'static RwLock< Vec< AtomEval<Self> > >;
}
}
macro_rules! impl_atom_eval_vec{ ($V:ty) => {
#[doc = concat!(
"The atomic evaluation vector for value type `", stringify!($V), "`"
) ]
impl crate::atom::sealed::AtomEvalVec for $V {
fn get() -> &'static
RwLock< Vec< crate::atom::AtomEval<$V> > > {
pub(crate) static ATOM_EVAL_VEC :
RwLock< Vec< crate::atom::AtomEval<$V> > > =
RwLock::new( Vec::new() );
&ATOM_EVAL_VEC
}
}
} }
pub(crate) use impl_atom_eval_vec;
pub fn register_atom<V>( atom_eval : AtomEval<V> ) -> IndexT
where
V : AtomEvalVecPublic ,
{ let rw_lock : &RwLock< Vec< AtomEval<V> > > = sealed::AtomEvalVec::get();
let atom_id : IndexT;
let atom_id_too_large : bool;
{ let write_lock = rw_lock.write();
assert!( write_lock.is_ok() );
let mut atom_eval_vec = write_lock.unwrap();
let atom_id_usize = atom_eval_vec.len();
atom_id_too_large = (IndexT::MAX as usize) < atom_id_usize;
atom_id = atom_eval_vec.len() as IndexT;
atom_eval_vec.push( atom_eval );
}
assert!( ! atom_id_too_large );
atom_id
}
fn record_call_atom<V>(
tape : &mut Tape<V> ,
forward_depend : AtomForwardDepend ,
adomain : Vec< AD<V> > ,
range_zero : Vec<V> ,
atom_id : IndexT ,
call_info : IndexT ,
trace : bool ,
) -> Vec< AD<V> >
where
V : Clone ,
{ debug_assert!( tape.recording );
let call_n_arg = adomain.len();
let call_n_res = range_zero.len();
let mut arange : Vec< AD<V> > = ad_from_vector(range_zero);
let is_var_arg : Vec<bool> = adomain.iter().map(
|adomain_j| (*adomain_j).tape_id == tape.tape_id
).collect();
let is_var_res = forward_depend(&is_var_arg, call_info, trace);
let mut n_var_res = 0;
for i in 0 .. call_n_res {
if is_var_res[i] {
arange[i].tape_id = tape.tape_id;
arange[i].var_index = tape.n_var + n_var_res;
n_var_res += 1;
}
}
if n_var_res > 0 {
tape.id_all.push( CALL_OP );
tape.op2arg.push( tape.arg_all.len() as IndexT );
tape.arg_all.push( atom_id ); tape.arg_all.push( call_info ); tape.arg_all.push( call_n_arg as IndexT ); tape.arg_all.push( call_n_res as IndexT ); tape.arg_all.push( tape.flag_all.len() as IndexT ); for j in 0 .. call_n_arg {
let index = if is_var_arg[j] {
adomain[j].var_index
} else {
let con_index = tape.con_all.len();
tape.con_all.push( adomain[j].value.clone() );
con_index
};
tape.arg_all.push( index as IndexT ); }
tape.flag_all.push( trace ); for j in 0 .. call_n_arg {
tape.flag_all.push( is_var_arg[j] ); }
for i in 0 .. call_n_res {
tape.flag_all.push( is_var_res[i] ); }
tape.n_var += n_var_res;
for _i in 0 .. (n_var_res - 1) {
tape.id_all.push( CALL_RES_OP );
tape.op2arg.push( tape.arg_all.len() as IndexT );
}
}
arange
}
pub fn call_atom<V>(
adomain : Vec< AD<V> > ,
atom_id : IndexT ,
call_info : IndexT ,
trace : bool ,
) -> Vec< AD<V> >
where
V : Clone + From<f32> + ThisThreadTapePublic + AtomEvalVecPublic ,
{
let local_key : &LocalKey< RefCell< Tape<V> > > = ThisThreadTape::get();
let recording : bool = local_key.with_borrow( |tape| tape.recording );
let rw_lock : &RwLock< Vec< AtomEval<V> > > = sealed::AtomEvalVec::get();
let forward_zero : AtomForwardZeroValue<V>;
let forward_depend : AtomForwardDepend;
{ 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 as usize];
forward_zero = atom_eval.forward_zero_value.clone();
forward_depend = atom_eval.forward_depend.clone();
}
let mut domain_zero : Vec<&V> = Vec::with_capacity( adomain.len() );
for j in 0 .. adomain.len() {
domain_zero.push( &adomain[j].value );
}
let range_zero = forward_zero( &domain_zero, call_info, trace );
let arange : Vec< AD<V> >;
if ! recording {
arange = ad_from_vector(range_zero);
} else {
arange = local_key.with_borrow_mut( |tape| record_call_atom::<V>(
tape,
forward_depend,
adomain,
range_zero,
atom_id,
call_info,
trace,
) );
}
arange
}