Stan  1.3
probability, sampling & optimization
 All Classes Namespaces Files Functions Variables Typedefs Enumerator Friends Macros Pages
dist.hpp
Go to the documentation of this file.
1 #ifndef __STAN__AGRAD__REV__MATRIX__DIST_HPP__
2 #define __STAN__AGRAD__REV__MATRIX__DIST_HPP__
3 
4 #include <vector>
9 #include <stan/agrad/rev/var.hpp>
10 #include <stan/agrad/rev/vari.hpp>
11 #include <stan/agrad/rev/sqrt.hpp>
13 
14 namespace stan {
15  namespace agrad {
16  namespace {
17  class squared_dist_vv_vari : public vari {
18  protected:
19  vari** v1_;
20  vari** v2_;
21  size_t length_;
22 
23  template<int R1,int C1,int R2,int C2>
24  inline static double var_squared_dist(const Eigen::Matrix<var,R1,C1> &v1,
25  const Eigen::Matrix<var,R2,C2> &v2) {
26  double result = 0;
27  for (size_t i = 0; i < v1.size(); i++) {
28  double diff = v1[i].vi_->val_ - v2[i].vi_->val_;
29  result += diff*diff;
30  }
31  return result;
32  }
33  public:
34  template<int R1,int C1,int R2,int C2>
35  squared_dist_vv_vari(const Eigen::Matrix<var,R1,C1> &v1,
36  const Eigen::Matrix<var,R2,C2> &v2)
37  : vari(var_squared_dist(v1, v2)), length_(v1.size())
38  {
39  v1_ = (vari**)memalloc_.alloc(length_*sizeof(vari*));
40  for (size_t i = 0; i < length_; i++)
41  v1_[i] = v1[i].vi_;
42 
43  v2_ = (vari**)memalloc_.alloc(length_*sizeof(vari*));
44  for (size_t i = 0; i < length_; i++)
45  v2_[i] = v2[i].vi_;
46  }
47  virtual void chain() {
48  for (size_t i = 0; i < length_; i++) {
49  double di = 2 * adj_ * (v1_[i]->val_ - v2_[i]->val_);
50  v1_[i]->adj_ += di;
51  v2_[i]->adj_ -= di;
52  }
53  }
54  };
55  class squared_dist_vd_vari : public vari {
56  protected:
57  vari** v1_;
58  double* v2_;
59  size_t length_;
60 
61  template<int R1,int C1,int R2,int C2>
62  inline static double var_squared_dist(const Eigen::Matrix<var,R1,C1> &v1,
63  const Eigen::Matrix<double,R2,C2> &v2) {
64  double result = 0;
65  for (size_t i = 0; i < v1.size(); i++) {
66  double diff = v1[i].vi_->val_ - v2[i];
67  result += diff*diff;
68  }
69  return result;
70  }
71  public:
72  template<int R1,int C1,int R2,int C2>
73  squared_dist_vd_vari(const Eigen::Matrix<var,R1,C1> &v1,
74  const Eigen::Matrix<double,R2,C2> &v2)
75  : vari(var_squared_dist(v1, v2)), length_(v1.size())
76  {
77  v1_ = (vari**)memalloc_.alloc(length_*sizeof(vari*));
78  for (size_t i = 0; i < length_; i++)
79  v1_[i] = v1[i].vi_;
80 
81  v2_ = (double*)memalloc_.alloc(length_*sizeof(double));
82  for (size_t i = 0; i < length_; i++)
83  v2_[i] = v2[i];
84  }
85  virtual void chain() {
86  for (size_t i = 0; i < length_; i++) {
87  v1_[i]->adj_ += 2 * adj_ * (v1_[i]->val_ - v2_[i]);
88  }
89  }
90  };
91  }
92 
93  template<int R1,int C1,int R2, int C2>
94  inline var squared_dist(const Eigen::Matrix<var, R1, C1>& v1,
95  const Eigen::Matrix<var, R2, C2>& v2) {
96  stan::math::validate_vector(v1,"squared_dist");
97  stan::math::validate_vector(v2,"squared_dist");
98  stan::math::validate_matching_sizes(v1,v2,"squared_dist");
99  return var(new squared_dist_vv_vari(v1,v2));
100  }
101  template<int R1,int C1,int R2, int C2>
102  inline var squared_dist(const Eigen::Matrix<var, R1, C1>& v1,
103  const Eigen::Matrix<double, R2, C2>& v2) {
104  stan::math::validate_vector(v1,"squared_dist");
105  stan::math::validate_vector(v2,"squared_dist");
106  stan::math::validate_matching_sizes(v1,v2,"squared_dist");
107  return var(new squared_dist_vd_vari(v1,v2));
108  }
109  template<int R1,int C1,int R2, int C2>
110  inline var squared_dist(const Eigen::Matrix<double, R1, C1>& v1,
111  const Eigen::Matrix<var, R2, C2>& v2) {
112  stan::math::validate_vector(v1,"squared_dist");
113  stan::math::validate_vector(v2,"squared_dist");
114  stan::math::validate_matching_sizes(v1,v2,"squared_dist");
115  return var(new squared_dist_vd_vari(v2,v1));
116  }
117 
118  template<int R1,int C1,int R2, int C2>
119  inline var dist(const Eigen::Matrix<var, R1, C1>& v1,
120  const Eigen::Matrix<var, R2, C2>& v2) {
121  stan::math::validate_vector(v1,"dist");
122  stan::math::validate_vector(v2,"dist");
124  return sqrt(var(new squared_dist_vv_vari(v1,v2)));
125  }
126  template<int R1,int C1,int R2, int C2>
127  inline var dist(const Eigen::Matrix<var, R1, C1>& v1,
128  const Eigen::Matrix<double, R2, C2>& v2) {
129  stan::math::validate_vector(v1,"dist");
130  stan::math::validate_vector(v2,"dist");
132  return sqrt(var(new squared_dist_vd_vari(v1,v2)));
133  }
134  template<int R1,int C1,int R2, int C2>
135  inline var dist(const Eigen::Matrix<double, R1, C1>& v1,
136  const Eigen::Matrix<var, R2, C2>& v2) {
137  stan::math::validate_vector(v1,"dist");
138  stan::math::validate_vector(v2,"dist");
140  return sqrt(var(new squared_dist_vd_vari(v2,v1)));
141  }
142  }
143 }
144 #endif

     [ Stan Home Page ] © 2011–2013, Stan Development Team.