Automatic Differentiation
 
Loading...
Searching...
No Matches
squared_distance.hpp
Go to the documentation of this file.
1#ifndef STAN_MATH_REV_FUN_SQUARED_DISTANCE_HPP
2#define STAN_MATH_REV_FUN_SQUARED_DISTANCE_HPP
3
13#include <vector>
14
15namespace stan {
16namespace math {
17
21inline var squared_distance(const var& a, const var& b) {
22 check_finite("squared_distance", "a", a);
23 check_finite("squared_distance", "b", b);
24 double difference = a.val() - b.val();
25 return make_callback_vari(difference * difference,
26 [a, b](const auto& vi) mutable {
27 const double diff = 2.0 * (a.val() - b.val());
28 a.adj() += vi.adj_ * diff;
29 b.adj() -= vi.adj_ * diff;
30 });
31}
32
36inline var squared_distance(const var& a, double b) {
37 check_finite("squared_distance", "a", a);
38 check_finite("squared_distance", "b", b);
39 double difference = a.val() - b;
40 return make_callback_vari(difference * difference,
41 [a, b](const auto& vi) mutable {
42 a.adj() += vi.adj_ * 2.0 * (a.val() - b);
43 });
44}
45
49inline var squared_distance(double a, const var& b) {
50 return squared_distance(b, a);
51}
52
53namespace internal {
54
56 protected:
59 size_t length_;
60
61 public:
62 template <
63 typename EigVecVar1, typename EigVecVar2,
65 squared_distance_vv_vari(const EigVecVar1& v1, const EigVecVar2& v2)
68 .squaredNorm()),
69 length_(v1.size()) {
70 v1_ = reinterpret_cast<vari**>(
72 v2_ = reinterpret_cast<vari**>(
74 Eigen::Map<vector_vi>(v1_, length_) = v1.vi();
75 Eigen::Map<vector_vi>(v2_, length_) = v2.vi();
76 }
77
78 virtual void chain() {
79 Eigen::Map<vector_vi> v1_map(v1_, length_);
80 Eigen::Map<vector_vi> v2_map(v2_, length_);
81 vector_d di = 2 * adj_ * (v1_map.val() - v2_map.val());
82 v1_map.adj() += di;
83 v2_map.adj() -= di;
84 }
85};
86
88 protected:
90 double* v2_;
91 size_t length_;
92
93 public:
94 template <typename EigVecVar, typename EigVecArith,
97 squared_distance_vd_vari(const EigVecVar& v1, const EigVecArith& v2)
100 .squaredNorm()),
101 length_(v1.size()) {
102 v1_ = reinterpret_cast<vari**>(
104 v2_ = reinterpret_cast<double*>(
106 Eigen::Map<vector_vi>(v1_, length_) = v1.vi();
107 Eigen::Map<vector_d>(v2_, length_) = v2;
108 }
109
110 virtual void chain() {
111 Eigen::Map<vector_vi> v1_map(v1_, length_);
112 v1_map.adj()
113 += 2 * adj_ * (v1_map.val() - Eigen::Map<vector_d>(v2_, length_));
114 }
115};
116} // namespace internal
117
118template <
119 typename EigVecVar1, typename EigVecVar2,
121inline var squared_distance(const EigVecVar1& v1, const EigVecVar2& v2) {
122 check_matching_sizes("squared_distance", "v1", v1, "v2", v2);
123 return {new internal::squared_distance_vv_vari(to_ref(v1), to_ref(v2))};
124}
125
126template <typename EigVecVar, typename EigVecArith,
129inline var squared_distance(const EigVecVar& v1, const EigVecArith& v2) {
130 check_matching_sizes("squared_distance", "v1", v1, "v2", v2);
131 return {new internal::squared_distance_vd_vari(to_ref(v1), to_ref(v2))};
132}
133
134template <typename EigVecArith, typename EigVecVar,
137inline var squared_distance(const EigVecArith& v1, const EigVecVar& v2) {
138 check_matching_sizes("squared_distance", "v1", v1, "v2", v2);
139 return {new internal::squared_distance_vd_vari(to_ref(v2), to_ref(v1))};
140}
141
157template <typename T1, typename T2, require_all_vector_t<T1, T2>* = nullptr,
158 require_any_var_vector_t<T1, T2>* = nullptr>
159inline var squared_distance(const T1& A, const T2& B) {
160 check_matching_sizes("squared_distance", "A", A.val(), "B", B.val());
161 if (unlikely(A.size() == 0)) {
162 return var(0.0);
163 } else if constexpr (is_autodiff_v<T1> && is_autodiff_v<T2>) {
166 arena_t<Eigen::VectorXd> res_diff(arena_A.size());
167 double res_val = 0.0;
168 for (size_t i = 0; i < arena_A.size(); ++i) {
169 const double diff = arena_A.val().coeff(i) - arena_B.val().coeff(i);
170 res_diff.coeffRef(i) = diff;
171 res_val += diff * diff;
172 }
173 return var(make_callback_vari(
174 res_val, [arena_A, arena_B, res_diff](const auto& res) mutable {
175 const double res_adj = 2.0 * res.adj();
176 for (size_t i = 0; i < arena_A.size(); ++i) {
177 const double diff = res_adj * res_diff.coeff(i);
178 arena_A.adj().coeffRef(i) += diff;
179 arena_B.adj().coeffRef(i) -= diff;
180 }
181 }));
182 } else if constexpr (is_autodiff_v<T1>) {
185 arena_t<Eigen::VectorXd> res_diff(arena_A.size());
186 double res_val = 0.0;
187 for (size_t i = 0; i < arena_A.size(); ++i) {
188 const double diff = arena_A.val().coeff(i) - arena_B.coeff(i);
189 res_diff.coeffRef(i) = diff;
190 res_val += diff * diff;
191 }
192 return var(make_callback_vari(
193 res_val, [arena_A, arena_B, res_diff](const auto& res) mutable {
194 arena_A.adj() += 2.0 * res.adj() * res_diff;
195 }));
196 } else {
199 arena_t<Eigen::VectorXd> res_diff(arena_A.size());
200 double res_val = 0.0;
201 for (size_t i = 0; i < arena_A.size(); ++i) {
202 const double diff = arena_A.coeff(i) - arena_B.val().coeff(i);
203 res_diff.coeffRef(i) = diff;
204 res_val += diff * diff;
205 }
206 return var(make_callback_vari(
207 res_val, [arena_A, arena_B, res_diff](const auto& res) mutable {
208 arena_B.adj() -= 2.0 * res.adj() * res_diff;
209 }));
210 }
211}
212
213} // namespace math
214} // namespace stan
215#endif
squared_distance_vd_vari(const EigVecVar &v1, const EigVecArith &v2)
squared_distance_vv_vari(const EigVecVar1 &v1, const EigVecVar2 &v2)
void * alloc(size_t len)
Return a newly allocated block of memory of the appropriate size managed by the stack allocator.
#define unlikely(x)
require_all_t< container_type_check_base< is_eigen_vector, value_type_t, TypeCheck, Check >... > require_all_eigen_vector_vt
Require all of the types satisfy is_eigen_vector.
require_t< container_type_check_base< is_eigen_vector, value_type_t, TypeCheck, Check... > > require_eigen_vector_vt
Require type satisfies is_eigen_vector.
auto as_column_vector_or_scalar(T &&a)
as_column_vector_or_scalar of a kernel generator expression.
int64_t size(const T &m)
Returns the size (number of the elements) of a matrix_cl or var_value<matrix_cl<T>>.
Definition size.hpp:19
Eigen::Matrix< double, Eigen::Dynamic, 1 > vector_d
Type for (column) vector of double values.
Definition typedefs.hpp:24
T value_of(const fvar< T > &v)
Return the value of the specified variable.
Definition value_of.hpp:18
void check_matching_sizes(const char *function, const char *name1, const T_y1 &y1, const char *name2, const T_y2 &y2)
Check if two structures at the same size.
void check_finite(const char *function, const char *name, const T_y &y)
Return true if all values in y are finite.
var_value< double > var
Definition var.hpp:1187
auto squared_distance(const T_a &a, const T_b &b)
Returns the squared distance.
ref_type_t< T && > to_ref(T &&a)
This evaluates expensive Eigen expressions.
Definition to_ref.hpp:18
internal::callback_vari< plain_type_t< T >, F > * make_callback_vari(T &&value, F &&functor)
Creates a new vari with given value and a callback that implements the reverse pass (chain).
typename internal::arena_type_impl< std::decay_t< T > >::type arena_t
Determines a type that can be used in place of T that does any dynamic allocations on the AD stack.
The lgamma implementation in stan-math is based on either the reentrant safe lgamma_r implementation ...
static thread_local AutodiffStackStorage * instance_