Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
CoherentPadOp.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2020 The Eigen Team.
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11#ifndef EIGEN_COHERENT_PAD_OP_H
12#define EIGEN_COHERENT_PAD_OP_H
13
14#include "./InternalHeaderCheck.h"
15
16namespace Eigen {
17
18namespace internal {
19
20// Pads a vector with zeros to a given size.
21template <typename XprType, int SizeAtCompileTime_>
22struct CoherentPadOp;
23
24template <typename XprType, int SizeAtCompileTime_>
25struct traits<CoherentPadOp<XprType, SizeAtCompileTime_>> : public traits<XprType> {
26 typedef internal::remove_all_t<XprType> PlainXprType;
27 typedef typename internal::ref_selector<XprType>::type XprNested;
28 typedef std::remove_reference_t<XprNested> XprNested_;
29 enum : int {
30 IsRowMajor = traits<PlainXprType>::Flags & RowMajorBit,
31 SizeAtCompileTime = SizeAtCompileTime_,
32 RowsAtCompileTime = IsRowMajor ? 1 : SizeAtCompileTime,
33 ColsAtCompileTime = IsRowMajor ? SizeAtCompileTime : 1,
34 MaxRowsAtCompileTime = RowsAtCompileTime,
35 MaxColsAtCompileTime = ColsAtCompileTime,
36 Flags = traits<XprType>::Flags & ~NestByRefBit,
37 };
38};
39
40// Pads a vector with zeros to a given size.
41template <typename XprType, int SizeAtCompileTime_>
42struct CoherentPadOp : public dense_xpr_base<CoherentPadOp<XprType, SizeAtCompileTime_>>::type {
43 typedef typename internal::generic_xpr_base<CoherentPadOp<XprType, SizeAtCompileTime_>>::type Base;
44 EIGEN_GENERIC_PUBLIC_INTERFACE(CoherentPadOp)
45
46 using XprNested = typename traits<CoherentPadOp>::XprNested;
47 using XprNested_ = typename traits<CoherentPadOp>::XprNested_;
48 using NestedExpression = XprNested_;
49
50 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoherentPadOp() = delete;
51 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoherentPadOp(const CoherentPadOp&) = default;
52 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoherentPadOp(CoherentPadOp&& other) = default;
53
54 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoherentPadOp(const XprType& xpr, Index size) : m_xpr(xpr), m_size(size) {
55 static_assert(XprNested_::IsVectorAtCompileTime, "input type must be a vector");
56 }
57
58 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const XprNested_& nestedExpression() const { return m_xpr; }
59
60 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index size() const { return m_size.value(); }
61
62 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index rows() const {
63 return traits<CoherentPadOp>::IsRowMajor ? Index(1) : size();
64 }
65
66 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index cols() const {
67 return traits<CoherentPadOp>::IsRowMajor ? size() : Index(1);
68 }
69
70 private:
71 XprNested m_xpr;
72 const internal::variable_if_dynamic<Index, SizeAtCompileTime> m_size;
73};
74
75// Adapted from the Replicate evaluator.
76template <typename ArgType, int SizeAtCompileTime>
77struct unary_evaluator<CoherentPadOp<ArgType, SizeAtCompileTime>>
78 : evaluator_base<CoherentPadOp<ArgType, SizeAtCompileTime>> {
79 typedef CoherentPadOp<ArgType, SizeAtCompileTime> XprType;
80 typedef internal::remove_all_t<typename XprType::CoeffReturnType> CoeffReturnType;
81 typedef typename internal::nested_eval<ArgType, 1>::type ArgTypeNested;
82 typedef internal::remove_all_t<ArgTypeNested> ArgTypeNestedCleaned;
83
84 enum {
85 CoeffReadCost = evaluator<ArgTypeNestedCleaned>::CoeffReadCost,
86 LinearAccessMask = XprType::IsVectorAtCompileTime ? LinearAccessBit : 0,
87 Flags = evaluator<ArgTypeNestedCleaned>::Flags & (HereditaryBits | LinearAccessMask | RowMajorBit),
88 Alignment = evaluator<ArgTypeNestedCleaned>::Alignment
89 };
90
91 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE explicit unary_evaluator(const XprType& pad)
92 : m_arg(pad.nestedExpression()), m_argImpl(m_arg), m_size(pad.nestedExpression().size()) {}
93
94 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index row, Index col) const {
95 EIGEN_IF_CONSTEXPR (XprType::IsRowMajor) {
96 if (col < m_size.value()) {
97 return m_argImpl.coeff(1, col);
98 }
99 } else {
100 if (row < m_size.value()) {
101 return m_argImpl.coeff(row, 1);
102 }
103 }
104 return CoeffReturnType(0);
105 }
106
107 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index) const {
108 if (index < m_size.value()) {
109 return m_argImpl.coeff(index);
110 }
111 return CoeffReturnType(0);
112 }
113
114 template <int LoadMode, typename PacketType>
115 EIGEN_STRONG_INLINE PacketType packet(Index row, Index col) const {
116 // AutoDiff scalar's derivative must be a vector, which is enforced by static assert.
117 // Defer to linear access for simplicity.
118 EIGEN_IF_CONSTEXPR (XprType::IsRowMajor) {
119 return packet(col);
120 }
121 return packet(row);
122 }
123
124 template <int LoadMode, typename PacketType>
125 EIGEN_STRONG_INLINE PacketType packet(Index index) const {
126 constexpr int kPacketSize = unpacket_traits<PacketType>::size;
127 if (index + kPacketSize <= m_size.value()) {
128 return m_argImpl.template packet<LoadMode, PacketType>(index);
129 } else if (index < m_size.value()) {
130 // Partial packet.
131 EIGEN_ALIGN_TO_BOUNDARY(unpacket_traits<PacketType>::alignment)
132 std::remove_const_t<CoeffReturnType> values[kPacketSize];
133 const int partial = m_size.value() - index;
134 for (int i = 0; i < partial && i < kPacketSize; ++i) {
135 values[i] = m_argImpl.coeff(index + i);
136 }
137 for (int i = partial; i < kPacketSize; ++i) {
138 values[i] = CoeffReturnType(0);
139 }
140 return pload<PacketType>(values);
141 }
142 return pset1<PacketType>(CoeffReturnType(0));
143 }
144
145 protected:
146 ArgTypeNested m_arg;
147 evaluator<ArgTypeNestedCleaned> m_argImpl;
148 const variable_if_dynamic<Index, ArgTypeNestedCleaned::SizeAtCompileTime> m_size;
149};
150
151} // namespace internal
152
153} // namespace Eigen
154
155#endif // EIGEN_COHERENT_PAD_OP_H
constexpr unsigned int LinearAccessBit
constexpr unsigned int RowMajorBit
Namespace containing all symbols from the Eigen library.