1 #ifndef __STAN__AGRAD__AGRAD_HPP__
2 #define __STAN__AGRAD__AGRAD_HPP__
20 struct var_allocator {
23 inline void* alloc(
size_t nbytes) {
26 inline void recover() {
34 #ifdef AGRAD_THREAD_SAFE
37 var_allocator* allocator_;
85 allocator_->var_stack_.push_back(
this);
105 static inline void*
operator new(
size_t nbytes) {
107 allocator_ =
new var_allocator();
108 return allocator_->alloc(nbytes);
115 return allocator_->recover();
138 std::vector<vari*>::iterator it = allocator_->var_stack_.end();
139 std::vector<vari*>::iterator begin = allocator_->var_stack_.begin();
141 for (; (it >= begin) && (*it != vi); --it)
145 for (; it >= begin; --it)
153 class op_v_vari :
public vari {
157 op_v_vari(
double f, vari* avi) :
163 class op_vv_vari :
public vari {
168 op_vv_vari(
double f, vari* avi, vari* bvi):
175 class op_vd_vari :
public vari {
180 op_vd_vari(
double f, vari* avi,
double b) :
187 class op_dv_vari :
public vari {
192 op_dv_vari(
double f,
double a, vari* bvi) :
199 class op_vvv_vari :
public vari {
205 op_vvv_vari(
double f, vari* avi, vari* bvi, vari* cvi) :
213 class op_vvd_vari :
public vari {
219 op_vvd_vari(
double f, vari* avi, vari* bvi,
double c) :
227 class op_vdv_vari :
public vari {
233 op_vdv_vari(
double f, vari* avi,
double b, vari* cvi) :
241 class op_vdd_vari :
public vari {
247 op_vdd_vari(
double f, vari* avi,
double b,
double c) :
255 class op_dvv_vari :
public vari {
261 op_dvv_vari(
double f,
double a, vari* bvi, vari* cvi) :
269 class op_dvd_vari :
public vari {
275 op_dvd_vari(
double f,
double a, vari* bvi,
double c) :
283 class op_ddv_vari :
public vari {
289 op_ddv_vari(
double f,
double a,
double b, vari* cvi) :
297 class neg_vari :
public op_v_vari {
299 neg_vari(vari* avi) :
300 op_v_vari(-(avi->val_), avi) {
308 class add_vv_vari :
public op_vv_vari {
310 add_vv_vari(vari* avi, vari* bvi) :
311 op_vv_vari(avi->val_ + bvi->val_, avi, bvi) {
319 class add_vd_vari :
public op_vd_vari {
321 add_vd_vari(vari* avi,
double b) :
322 op_vd_vari(avi->val_ + b, avi, b) {
329 class increment_vari :
public op_v_vari {
331 increment_vari(vari* avi) :
332 op_v_vari(avi->val_ + 1.0, avi) {
339 class decrement_vari :
public op_v_vari {
341 decrement_vari(vari* avi) :
342 op_v_vari(avi->val_ - 1.0, avi) {
349 class subtract_vv_vari :
public op_vv_vari {
351 subtract_vv_vari(vari* avi, vari* bvi) :
352 op_vv_vari(avi->val_ - bvi->val_, avi, bvi) {
360 class subtract_vd_vari :
public op_vd_vari {
362 subtract_vd_vari(vari* avi,
double b) :
363 op_vd_vari(avi->val_ - b, avi, b) {
370 class subtract_dv_vari :
public op_dv_vari {
372 subtract_dv_vari(
double a, vari* bvi) :
373 op_dv_vari(a - bvi->val_, a, bvi) {
380 class multiply_vv_vari :
public op_vv_vari {
382 multiply_vv_vari(vari* avi, vari* bvi) :
383 op_vv_vari(avi->val_ * bvi->val_, avi, bvi) {
391 class multiply_vd_vari :
public op_vd_vari {
393 multiply_vd_vari(vari* avi,
double b) :
394 op_vd_vari(avi->val_ * b, avi, b) {
402 class divide_vv_vari :
public op_vv_vari {
404 divide_vv_vari(vari* avi, vari* bvi) :
405 op_vv_vari(avi->val_ / bvi->val_, avi, bvi) {
413 class divide_vd_vari :
public op_vd_vari {
415 divide_vd_vari(vari* avi,
double b) :
416 op_vd_vari(avi->val_ / b, avi, b) {
423 class divide_dv_vari :
public op_dv_vari {
425 divide_dv_vari(
double a, vari* bvi) :
426 op_dv_vari(a / bvi->val_, a, bvi) {
433 class exp_vari :
public op_v_vari {
435 exp_vari(vari* avi) :
436 op_v_vari(std::
exp(avi->val_),avi) {
439 avi_->adj_ += adj_ * val_;
443 class log_vari :
public op_v_vari {
445 log_vari(vari* avi) :
446 op_v_vari(std::
log(avi->val_),avi) {
455 class log10_vari :
public op_v_vari {
458 log10_vari(vari* avi) :
459 op_v_vari(std::
log10(avi->val_),avi),
467 class sqrt_vari :
public op_v_vari {
469 sqrt_vari(vari* avi) :
470 op_v_vari(std::
sqrt(avi->val_),avi) {
473 avi_->adj_ += adj_ / (2.0 * val_);
477 class pow_vv_vari :
public op_vv_vari {
479 pow_vv_vari(vari* avi, vari* bvi) :
480 op_vv_vari(std::
pow(avi->val_,bvi->val_),avi,bvi) {
483 if (
avi_->val_ == 0.0)
return;
489 class pow_vd_vari :
public op_vd_vari {
491 pow_vd_vari(vari* avi,
double b) :
492 op_vd_vari(std::
pow(avi->val_,b),avi,b) {
495 if (
avi_->val_ == 0.0)
return;
500 class pow_dv_vari :
public op_dv_vari {
502 pow_dv_vari(
double a, vari* bvi) :
503 op_dv_vari(std::
pow(a,bvi->val_),a,bvi) {
506 if (
ad_ == 0.0)
return;
511 class cos_vari :
public op_v_vari {
513 cos_vari(vari* avi) :
514 op_v_vari(std::
cos(avi->val_),avi) {
521 class sin_vari :
public op_v_vari {
523 sin_vari(vari* avi) :
524 op_v_vari(std::
sin(avi->val_),avi) {
531 class tan_vari :
public op_v_vari {
533 tan_vari(vari* avi) :
534 op_v_vari(std::
tan(avi->val_),avi) {
537 avi_->adj_ += adj_ * (1.0 + val_ * val_);
541 class acos_vari :
public op_v_vari {
543 acos_vari(vari* avi) :
544 op_v_vari(std::
acos(avi->val_),avi) {
551 class asin_vari :
public op_v_vari {
553 asin_vari(vari* avi) :
554 op_v_vari(std::
asin(avi->val_),avi) {
561 class atan_vari :
public op_v_vari {
563 atan_vari(vari* avi) :
564 op_v_vari(std::
atan(avi->val_),avi) {
571 class atan2_vv_vari :
public op_vv_vari {
573 atan2_vv_vari(vari* avi, vari* bvi) :
574 op_vv_vari(std::
atan2(avi->val_,bvi->val_),avi,bvi) {
577 double a_sq_plus_b_sq = (
avi_->val_ *
avi_->val_) + (
bvi_->val_ *
bvi_->val_);
578 avi_->adj_ +=
bvi_->val_ / a_sq_plus_b_sq;
579 bvi_->adj_ -=
avi_->val_ / a_sq_plus_b_sq;
583 class atan2_vd_vari :
public op_vd_vari {
585 atan2_vd_vari(vari* avi,
double b) :
586 op_vd_vari(std::
atan2(avi->val_,b),avi,b) {
590 avi_->adj_ += bd_ / a_sq_plus_b_sq;
594 class atan2_dv_vari :
public op_dv_vari {
596 atan2_dv_vari(
double a, vari* bvi) :
597 op_dv_vari(std::
atan2(a,bvi->val_),a,bvi) {
601 bvi_->adj_ -=
ad_ / a_sq_plus_b_sq;
605 class cosh_vari :
public op_v_vari {
607 cosh_vari(vari* avi) :
608 op_v_vari(std::
cosh(avi->val_),avi) {
615 class sinh_vari :
public op_v_vari {
617 sinh_vari(vari* avi) :
618 op_v_vari(std::
sinh(avi->val_),avi) {
625 class tanh_vari :
public op_v_vari {
627 tanh_vari(vari* avi) :
628 op_v_vari(std::
tanh(avi->val_),avi) {
637 class floor_vari :
public vari {
639 floor_vari(vari* avi) :
640 vari(std::
floor(avi->val_)) {
644 class ceil_vari :
public vari {
646 ceil_vari(vari* avi) :
647 vari(std::
ceil(avi->val_)) {
651 class fmod_vv_vari :
public op_vv_vari {
653 fmod_vv_vari(vari* avi, vari* bvi) :
654 op_vv_vari(std::
fmod(avi->val_,bvi->val_),avi,bvi) {
658 bvi_->adj_ -= adj_ *
static_cast<int>(
avi_->val_ /
bvi_->val_);
662 class fmod_vd_vari :
public op_v_vari {
664 fmod_vd_vari(vari* avi,
double b) :
665 op_v_vari(std::
fmod(avi->val_,b),avi) {
672 class fmod_dv_vari :
public op_dv_vari {
674 fmod_dv_vari(
double a, vari* bvi) :
675 op_dv_vari(std::
fmod(a,bvi->val_),a,bvi) {
678 int d =
static_cast<int>(
ad_ /
bvi_->val_);
679 bvi_->adj_ -= adj_ * d;
742 vi_(new
vari(static_cast<double>(b))) {
752 vi_(new
vari(static_cast<double>(c))) {
762 vi_(new
vari(static_cast<double>(n))) {
772 vi_(new
vari(static_cast<double>(n))) {
782 vi_(new
vari(static_cast<double>(n))) {
792 vi_(new
vari(static_cast<double>(n))) {
802 vi_(new
vari(static_cast<double>(n))) {
811 var(
unsigned long int n) :
812 vi_(new
vari(static_cast<double>(n))) {
822 vi_(new
vari(static_cast<double>(x))) {
841 vi_(new
vari(static_cast<double>(x))) {
849 inline double val()
const {
866 std::vector<double>& g) {
869 for (
size_t i = 0U; i < x.size(); ++i)
908 vi_ =
new add_vv_vari(
vi_,b.vi_);
923 vi_ =
new add_vd_vari(
vi_,b);
939 vi_ =
new subtract_vv_vari(
vi_,b.vi_);
955 vi_ =
new subtract_vd_vari(
vi_,b);
971 vi_ =
new multiply_vv_vari(
vi_,b.vi_);
987 vi_ =
new multiply_vd_vari(
vi_,b);
1002 vi_ =
new divide_vv_vari(
vi_,b.vi_);
1018 vi_ =
new divide_vd_vari(
vi_,b);
1036 return a.val() == b.val();
1049 return a.val() == b;
1061 return a == b.val();
1073 return a.val() != b.val();
1086 return a.val() != b;
1099 return a != b.val();
1110 return a.val() < b.val();
1145 return a.val() > b.val();
1182 return a.val() <= b.val();
1195 return a.val() <= b;
1208 return a <= b.val();
1221 return a.val() >= b.val();
1234 return a.val() >= b;
1247 return a >= b.val();
1299 return var(
new neg_vari(a.vi_));
1316 return var(
new add_vv_vari(a.vi_,b.vi_));
1332 return var(
new add_vd_vari(a.vi_,b));
1347 return var(
new add_vd_vari(b.vi_,a));
1365 return var(
new subtract_vv_vari(a.vi_,b.vi_));
1380 return var(
new subtract_vd_vari(a.vi_,b));
1395 return var(
new subtract_dv_vari(a,b.vi_));
1412 return var(
new multiply_vv_vari(a.vi_,b.vi_));
1427 return var(
new multiply_vd_vari(a.vi_,b));
1442 return var(
new multiply_vd_vari(b.vi_,a));
1460 return var(
new divide_vv_vari(a.vi_,b.vi_));
1475 return var(
new divide_vd_vari(a.vi_,b));
1490 return var(
new divide_dv_vari(a,b.vi_));
1503 a.vi_ =
new increment_vari(a.vi_);
1522 a.vi_ =
new increment_vari(a.vi_);
1540 a.vi_ =
new decrement_vari(a.vi_);
1559 a.vi_ =
new decrement_vari(a.vi_);
1572 return var(
new exp_vari(a.vi_));
1586 return var(
new log_vari(a.vi_));
1600 return var(
new log10_vari(a.vi_));
1617 return var(
new sqrt_vari(a.vi_));
1634 return var(
new pow_vv_vari(base.vi_,exponent.vi_));
1650 return var(
new pow_vd_vari(base.vi_,exponent));
1666 return var(
new pow_dv_vari(base,exponent.vi_));
1683 return var(
new cos_vari(a.vi_));
1697 return var(
new sin_vari(a.vi_));
1711 return var(
new tan_vari(a.vi_));
1726 return var(
new acos_vari(a.vi_));
1741 return var(
new asin_vari(a.vi_));
1756 return var(
new atan_vari(a.vi_));
1775 return var(
new atan2_vv_vari(a.vi_,b.vi_));
1791 return var(
new atan2_vd_vari(a.vi_,b));
1807 return var(
new atan2_dv_vari(a,b.vi_));
1823 return var(
new cosh_vari(a.vi_));
1837 return var(
new sinh_vari(a.vi_));
1851 return var(
new tanh_vari(a.vi_));
1878 return var(
new neg_vari(a.vi_));
1901 return var(
new floor_vari(a.vi_));
1923 return var(
new ceil_vari(a.vi_));
1944 return var(
new fmod_vv_vari(a.vi_,b.vi_));
1961 return var(
new fmod_vd_vari(a.vi_,b));
1978 return var(
new fmod_dv_vari(a,b.vi_));
2003 return var(
new neg_vari(a.vi_));