1#ifndef STAN_MATH_PRIM_FUN_HOLDER_HPP
2#define STAN_MATH_PRIM_FUN_HOLDER_HPP
69template <
class ArgType,
typename... Ptrs>
78template <
class ArgType,
typename... Ptrs>
83 typedef typename ArgType::Scalar
Scalar;
87 Flags = (ArgType::Flags
88 & (RowMajorBit | LvalueBit | LinearAccessBit | DirectAccessBit
89 | PacketAccessBit | NoPreferredStorageOrderBit))
91 RowsAtCompileTime = ArgType::RowsAtCompileTime,
92 ColsAtCompileTime = ArgType::ColsAtCompileTime,
93 MaxRowsAtCompileTime = ArgType::MaxRowsAtCompileTime,
94 MaxColsAtCompileTime = ArgType::MaxColsAtCompileTime,
95 InnerStrideAtCompileTime = ArgType::InnerStrideAtCompileTime,
96 OuterStrideAtCompileTime = ArgType::OuterStrideAtCompileTime
107template <
typename ArgType,
typename... Ptrs>
123 typename F,
typename... Args,
137template <
typename F,
typename... Args,
147template <
typename ArgType,
typename... Ptrs>
149 :
public Eigen::internal::dense_xpr_base<Holder<ArgType, Ptrs...>>::type {
151 typedef typename Eigen::internal::ref_selector<
Holder<ArgType, Ptrs...>>::type
153 typename Eigen::internal::ref_selector<ArgType>::non_const_type
m_arg;
176 return m_arg.coeffRef(index);
184 template <
typename T, require_eigen_t<T>* =
nullptr>
198 m_arg = std::move(other.m_arg);
203template <
typename T, require_holder_t<T>* =
nullptr>
207template <
typename T, require_holder_t<T>* =
nullptr>
212template <
typename T,
typename Other, require_holder_t<T>* =
nullptr,
213 require_holder_t<Other>* =
nullptr>
216 [](
auto&&
arg,
auto&& other_) {
217 return arg - std::forward<decltype(other_)>(other_);
219 std::forward<T>(h).m_arg, std::forward<Other>(other));
221template <
typename T,
typename Other, require_holder_t<T>* =
nullptr,
222 require_holder_t<Other>* =
nullptr>
225 [](
auto&&
arg,
auto&& other_) {
226 return arg + std::forward<decltype(other_)>(other_);
228 std::forward<T>(h).m_arg, std::forward<Other>(other));
230template <
typename T,
typename Other, require_holder_t<T>* =
nullptr,
231 require_holder_t<Other>* =
nullptr>
234 [](
auto&&
arg,
auto&& other_) {
235 return arg * std::forward<decltype(other_)>(other_);
237 std::forward<T>(h).m_arg, std::forward<Other>(other));
239template <
typename T,
typename Other, require_holder_t<T>* =
nullptr,
240 require_holder_t<Other>* =
nullptr>
243 [](
auto&&
arg,
auto&& other_) {
244 return arg / std::forward<decltype(other_)>(other_);
246 std::forward<T>(h).m_arg, std::forward<Other>(other));
255template <
typename ArgType,
typename... Ptrs>
256struct evaluator<
stan::math::Holder<ArgType, Ptrs...>>
257 : evaluator_base<stan::math::Holder<ArgType, Ptrs...>> {
264 IsRowMajor = XprType::IsRowMajor,
265 IsColMajor = !IsRowMajor,
266 IsVectorAtCompileTime = XprType::IsVectorAtCompileTime,
267 RowsAtCompileTime = XprType::RowsAtCompileTime,
268 ColsAtCompileTime = XprType::ColsAtCompileTime,
270 CoeffReadCost = evaluator<ArgTypeNestedCleaned>::CoeffReadCost,
271 Flags = evaluator<ArgTypeNestedCleaned>::Flags,
272 Alignment = evaluator<ArgTypeNestedCleaned>::Alignment,
276 OuterStrideAtCompileTime
277 = IsVectorAtCompileTime
279 : (IsRowMajor ? ColsAtCompileTime : RowsAtCompileTime)
286 : m_argImpl(
std::forward<
XprType>(xpr).m_arg) {}
291 return m_argImpl.coeff(row, col);
293 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType
295 return m_argImpl.coeff(index);
299 return m_argImpl.coeffRef(row, col);
302 return m_argImpl.coeffRef(index);
305 template <
int LoadMode,
typename PacketType>
306 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType
packet(Index row,
308 return m_argImpl.template packet<LoadMode, PacketType>(row, col);
310 template <
int LoadMode,
typename PacketType>
311 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType
packet(Index index)
const {
312 return m_argImpl.template packet<LoadMode, PacketType>(index);
315 template <
int StoreMode,
typename PacketType>
316 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void writePacket(Index row, Index col,
317 const PacketType& x) {
318 return m_argImpl.template writePacket<StoreMode, PacketType>(row, col, x);
320 template <
int StoreMode,
typename PacketType>
321 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void writePacket(Index index,
322 const PacketType& x) {
323 return m_argImpl.template writePacket<StoreMode, PacketType>(index, x);
326 template <
int LoadMode,
typename PacketType>
327 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType
329 return m_argImpl.template packetSegment<LoadMode, PacketType>(row, col,
333 template <
int LoadMode,
typename PacketType>
334 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType
336 return m_argImpl.template packetSegment<LoadMode, PacketType>(index, begin,
340 template <
int StoreMode,
typename PacketType>
342 Index row, Index col,
const PacketType& x, Index begin, Index count) {
343 return m_argImpl.template writePacketSegment<StoreMode, PacketType>(
344 row, col, x, begin, count);
347 template <
int StoreMode,
typename PacketType>
349 Index index,
const PacketType& x, Index begin, Index count) {
350 return m_argImpl.template writePacketSegment<StoreMode, PacketType>(
351 index, x, begin, count);
371template <
typename T,
typename... Ptrs,
372 std::enable_if_t<
sizeof...(Ptrs) >= 1>* =
nullptr>
374 return Holder<T, Ptrs...>(std::forward<T>(
arg), pointers...);
379 if constexpr (std::is_rvalue_reference<T&&>::value) {
380 return std::decay_t<T>(std::forward<T>(
arg));
382 return std::forward<T>(
arg);
400 return std::make_tuple();
403 std::enable_if_t<!(Eigen::internal::traits<std::decay_t<T>>::Flags
404 & Eigen::NestByRefBit)>* =
nullptr>
407 return std::make_tuple();
420template <
typename T, require_t<std::is_rvalue_reference<T&&>>* =
nullptr,
422 static_cast<
bool>(Eigen::
internal::traits<std::decay_t<T>>::Flags&
423 Eigen::NestByRefBit)>* =
nullptr>
425 res =
new T(std::move(a));
426 return std::make_tuple(res);
428template <
typename T, require_t<std::is_rvalue_reference<T&&>>* =
nullptr,
429 require_not_eigen_t<T>* =
nullptr>
430inline auto holder_handle_element(T&& a, T*& res) {
431 res =
new T(std::move(a));
432 return std::make_tuple(res);
446template <
typename T, std::size_t... Is,
typename... Args>
448 T&& expr, std::index_sequence<Is...>,
const std::tuple<Args*...>& ptrs) {
449 return holder(std::forward<T>(expr), std::get<Is>(ptrs)...);
461template <
typename F, std::size_t... Is,
typename... Args>
464 std::tuple<std::remove_reference_t<Args>*...> res;
465 auto ptrs = std::tuple_cat(
468 std::forward<F>(func)(*std::get<Is>(res)...),
469 std::make_index_sequence<std::tuple_size<
decltype(ptrs)>::value>(), ptrs);
486template <
typename F,
typename... Args,
490 return std::forward<F>(func)(std::forward<Args>(args)...);
493 std::forward<F>(func), std::make_index_sequence<
sizeof...(Args)>(),
494 std::forward<Args>(args)...);
508template <
typename F,
typename... Args,
511 return std::forward<F>(func)(std::forward<Args>(args)...);
Eigen::internal::ref_selector< Holder< ArgType, Ptrs... > >::type Nested
Holder< ArgType, Ptrs... > & operator=(const Holder< ArgType, Ptrs... > &other)
Eigen::Index outerStride() const
Eigen::Index rows() const
Eigen::Index innerStride() const
Holder< ArgType, Ptrs... > & operator=(Holder< ArgType, Ptrs... > &&other)
Holder(const Holder< ArgType, Ptrs... > &)=default
Holder< ArgType, Ptrs... > & operator=(const T &other)
Assignment operator assigns expressions.
const auto & coeffRef(Eigen::Index index) const
Eigen::Index cols() const
std::tuple< std::unique_ptr< Ptrs >... > m_unique_ptrs
const auto * data() const
Eigen::internal::ref_selector< ArgType >::non_const_type m_arg
Holder(Holder< ArgType, Ptrs... > &&)=default
const auto & coeffRef(Eigen::Index row, Eigen::Index col) const
Holder(ArgType &&arg, Ptrs *... pointers)
auto col(T_x &&x, size_t j)
Return the specified column of the specified kernel generator expression using start-at-1 indexing.
auto row(T_x &&x, size_t j)
Return the specified row of the specified kernel generator expression using start-at-1 indexing.
require_not_t< is_plain_type< std::decay_t< T > > > require_not_plain_type_t
Require type does not satisfy is_plain_type.
require_t< is_plain_type< std::decay_t< T > > > require_plain_type_t
Require type satisfies is_plain_type.
(Expert) Numerical traits for algorithmic differentiation variables.
auto make_holder_impl_construct_object(T &&expr, std::index_sequence< Is... >, const std::tuple< Args *... > &ptrs)
Second step in implementation of construction holder from a functor.
auto holder_handle_element(T &a, T *&res)
Handles single element (moving rvalue non-expressions to heap) for construction of holder or holder_c...
auto make_holder_impl(F &&func, std::index_sequence< Is... >, Args &&... args)
Implementation of construction holder from a functor.
fvar< T > operator/(const fvar< T > &x1, const fvar< T > &x2)
Return the result of dividing the first argument by the second.
fvar< T > operator-(const fvar< T > &x1, const fvar< T > &x2)
Return the difference of the specified arguments.
fvar< T > operator*(const fvar< T > &x, const fvar< T > &y)
Return the product of the two arguments.
fvar< T > arg(const std::complex< fvar< T > > &z)
Return the phase angle of the complex argument.
auto make_holder(F &&func, Args &&... args)
Calls given function with given arguments.
fvar< T > operator+(const fvar< T > &x1, const fvar< T > &x2)
Return the sum of the specified forward mode addends.
Ptrs holder(T &&arg, Ptrs *... pointers)
require_t< is_holder< T > > require_holder_t
constexpr bool is_holder_v
constexpr bool is_var_matrix_v
std::enable_if_t< Check::value > require_t
If condition is true, template is enabled.
The lgamma implementation in stan-math is based on either the reentrant safe lgamma_r implementation ...
typename XprType::CoeffReturnType CoeffReturnType
evaluator< ArgTypeNestedCleaned > m_argImpl
void writePacket(Index row, Index col, const PacketType &x)
void writePacket(Index index, const PacketType &x)
Scalar & coeffRef(Index index)
CoeffReturnType coeff(Index index) const
typename XprType::Scalar Scalar
CoeffReturnType coeff(Index row, Index col) const
typename remove_all< ArgType >::type ArgTypeNestedCleaned
void writePacketSegment(Index row, Index col, const PacketType &x, Index begin, Index count)
void writePacketSegment(Index index, const PacketType &x, Index begin, Index count)
PacketType packet(Index index) const
PacketType packetSegment(Index row, Index col, Index begin, Index count) const
PacketType packetSegment(Index index, Index begin, Index count) const
PacketType packet(Index row, Index col) const
Scalar & coeffRef(Index row, Index col)
evaluator(const XprType &xpr)
ArgType::StorageKind StorageKind
traits< ArgType >::XprKind XprKind
ArgType::StorageIndex StorageIndex