Stan  1.3
probability, sampling & optimization
 All Classes Namespaces Files Functions Variables Typedefs Enumerator Friends Macros Pages
multiply.hpp
Go to the documentation of this file.
1 #ifndef __STAN__AGRAD__REV__MATRIX__MULTIPLY_HPP__
2 #define __STAN__AGRAD__REV__MATRIX__MULTIPLY_HPP__
3 
4 #include <vector>
5 #include <boost/math/tools/promotion.hpp>
11 #include <stan/agrad/rev/var.hpp>
16 
17 namespace stan {
18  namespace agrad {
19 
26  template <typename T1, typename T2>
27  inline
28  typename boost::math::tools::promote_args<T1,T2>::type
29  multiply(const T1& v, const T2& c) {
30  return v * c;
31  }
32 
39  template<typename T1,typename T2,int R2,int C2>
40  inline Eigen::Matrix<var,R2,C2> multiply(const T1& c,
41  const Eigen::Matrix<T2, R2, C2>& m) {
42  // FIXME: pull out to eliminate overpromotion of one side
43  // move to matrix.hpp w. promotion?
44  return to_var(m) * to_var(c);
45  }
46 
53  template<typename T1,int R1,int C1,typename T2>
54  inline Eigen::Matrix<var,R1,C1> multiply(const Eigen::Matrix<T1, R1, C1>& m,
55  const T2& c) {
56  return to_var(m) * to_var(c);
57  }
58 
69  template<int R1,int C1,int R2,int C2>
70  inline Eigen::Matrix<var,R1,C2> multiply(const Eigen::Matrix<var,R1,C1>& m1,
71  const Eigen::Matrix<var,R2,C2>& m2) {
72  stan::math::validate_multiplicable(m1,m2,"multiply");
73  Eigen::Matrix<var,R1,C2> result(m1.rows(),m2.cols());
74  for (int i = 0; i < m1.rows(); i++) {
75  typename Eigen::Matrix<var,R1,C1>::ConstRowXpr crow(m1.row(i));
76  for (int j = 0; j < m2.cols(); j++) {
77  typename Eigen::Matrix<var,R2,C2>::ConstColXpr ccol(m2.col(j));
78  if (j == 0) {
79  if (i == 0) {
80  result(i,j) = var(new dot_product_vv_vari(crow,ccol));
81  }
82  else {
83  dot_product_vv_vari *v2 = static_cast<dot_product_vv_vari*>(result(0,j).vi_);
84  result(i,j) = var(new dot_product_vv_vari(crow,ccol,NULL,v2));
85  }
86  }
87  else {
88  if (i == 0) {
89  dot_product_vv_vari *v1 = static_cast<dot_product_vv_vari*>(result(i,0).vi_);
90  result(i,j) = var(new dot_product_vv_vari(crow,ccol,v1));
91  }
92  else /* if (i != 0 && j != 0) */ {
93  dot_product_vv_vari *v1 = static_cast<dot_product_vv_vari*>(result(i,0).vi_);
94  dot_product_vv_vari *v2 = static_cast<dot_product_vv_vari*>(result(0,j).vi_);
95  result(i,j) = var(new dot_product_vv_vari(crow,ccol,v1,v2));
96  }
97  }
98  }
99  }
100  return result;
101  }
102 
113  template<int R1,int C1,int R2,int C2>
114  inline Eigen::Matrix<var,R1,C2> multiply(const Eigen::Matrix<double,R1,C1>& m1,
115  const Eigen::Matrix<var,R2,C2>& m2) {
116  stan::math::validate_multiplicable(m1,m2,"multiply");
117  Eigen::Matrix<var,R1,C2> result(m1.rows(),m2.cols());
118  for (int i = 0; i < m1.rows(); i++) {
119  typename Eigen::Matrix<double,R1,C1>::ConstRowXpr crow(m1.row(i));
120  for (int j = 0; j < m2.cols(); j++) {
121  typename Eigen::Matrix<var,R2,C2>::ConstColXpr ccol(m2.col(j));
122  // result(i,j) = dot_product(crow,ccol);
123  if (j == 0) {
124  if (i == 0) {
125  result(i,j) = var(new dot_product_vd_vari(ccol,crow));
126  }
127  else {
128  dot_product_vd_vari *v2 = static_cast<dot_product_vd_vari*>(result(0,j).vi_);
129  result(i,j) = var(new dot_product_vd_vari(ccol,crow,v2,NULL));
130  }
131  }
132  else {
133  if (i == 0) {
134  dot_product_vd_vari *v1 = static_cast<dot_product_vd_vari*>(result(i,0).vi_);
135  result(i,j) = var(new dot_product_vd_vari(ccol,crow,NULL,v1));
136  }
137  else /* if (i != 0 && j != 0) */ {
138  dot_product_vd_vari *v1 = static_cast<dot_product_vd_vari*>(result(i,0).vi_);
139  dot_product_vd_vari *v2 = static_cast<dot_product_vd_vari*>(result(0,j).vi_);
140  result(i,j) = var(new dot_product_vd_vari(ccol,crow,v2,v1));
141  }
142  }
143  }
144  }
145  return result;
146  }
147 
158  template<int R1,int C1,int R2,int C2>
159  inline Eigen::Matrix<var,R1,C2> multiply(const Eigen::Matrix<var,R1,C1>& m1,
160  const Eigen::Matrix<double,R2,C2>& m2) {
161  stan::math::validate_multiplicable(m1,m2,"multiply");
162  Eigen::Matrix<var,R1,C2> result(m1.rows(),m2.cols());
163  for (int i = 0; i < m1.rows(); i++) {
164  typename Eigen::Matrix<var,R1,C1>::ConstRowXpr crow(m1.row(i));
165  for (int j = 0; j < m2.cols(); j++) {
166  typename Eigen::Matrix<double,R2,C2>::ConstColXpr ccol(m2.col(j));
167  // result(i,j) = dot_product(crow,ccol);
168  if (j == 0) {
169  if (i == 0) {
170  result(i,j) = var(new dot_product_vd_vari(crow,ccol));
171  }
172  else {
173  dot_product_vd_vari *v2 = static_cast<dot_product_vd_vari*>(result(0,j).vi_);
174  result(i,j) = var(new dot_product_vd_vari(crow,ccol,NULL,v2));
175  }
176  }
177  else {
178  if (i == 0) {
179  dot_product_vd_vari *v1 = static_cast<dot_product_vd_vari*>(result(i,0).vi_);
180  result(i,j) = var(new dot_product_vd_vari(crow,ccol,v1,NULL));
181  }
182  else /* if (i != 0 && j != 0) */ {
183  dot_product_vd_vari *v1 = static_cast<dot_product_vd_vari*>(result(i,0).vi_);
184  dot_product_vd_vari *v2 = static_cast<dot_product_vd_vari*>(result(0,j).vi_);
185  result(i,j) = var(new dot_product_vd_vari(crow,ccol,v1,v2));
186  }
187  }
188  }
189  }
190  return result;
191  }
192 
202  template <int C1,int R2>
203  inline var multiply(const Eigen::Matrix<var, 1, C1>& rv,
204  const Eigen::Matrix<var, R2, 1>& v) {
205  if (rv.size() != v.size())
206  throw std::domain_error("row vector and vector must be same length in multiply");
207  return dot_product(rv, v);
208  }
218  template <int C1,int R2>
219  inline var multiply(const Eigen::Matrix<double, 1, C1>& rv,
220  const Eigen::Matrix<var, R2, 1>& v) {
221  stan::math::validate_multiplicable(rv,v,"multiply");
222  return dot_product(rv, v);
223  }
233  template <int C1,int R2>
234  inline var multiply(const Eigen::Matrix<var, 1, C1>& rv,
235  const Eigen::Matrix<double, R2, 1>& v) {
236  stan::math::validate_multiplicable(rv,v,"multiply");
237  return dot_product(rv, v);
238  }
239 
240  }
241 }
242 #endif

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