10 #include <initializer_list>
11 #include <type_traits>
31 : data_(const_cast<void*>(reinterpret_cast<const void*>(
data))),
35 std::fill(shape_.begin(), shape_.end(), 0);
36 std::fill(strides_.begin(), strides_.end(), 0);
37 std::copy(
shape.begin(),
shape.end(), shape_.begin());
38 calculateSizeAndStrides();
42 : data_(const_cast<void*>(reinterpret_cast<const void*>(
data))),
46 std::fill(shape_.begin(), shape_.end(), 0);
47 std::fill(strides_.begin(), strides_.end(), 0);
48 for (std::size_t i = 0; i <
MAX_DIMS; ++i)
56 calculateSizeAndStrides();
64 std::fill(shape_.begin(), shape_.end(), 0);
65 std::fill(strides_.begin(), strides_.end(), 0);
66 for (std::size_t i = 0; i <
MAX_DIMS; ++i)
74 calculateSizeAndStrides();
81 return reinterpret_cast<T*
>(data_);
85 return reinterpret_cast<const T*
>(data_);
118 template <
typename... Idxs>
121 return const_cast<T&
>(
static_cast<const NDArray&
>(*this)(idxs...));
124 template <
typename... Idxs>
127 return reinterpret_cast<const T*
>(data_)[computeOffset({idxs...})];
133 auto const_result =
static_cast<const NDArray&
>(*this)[idx];
135 const_cast<T*
>(const_result.data()),
136 const_result.shape());
142 std::fill(new_shape.begin(), new_shape.end(), 0);
143 std::copy(shape_.begin() + 1, shape_.end(), new_shape.begin());
146 reinterpret_cast<const T*
>(data_) + idx * strides_[0],
151 template <
typename U = T>
152 operator typename std::enable_if_t<std::is_same_v<U, std::remove_const_t<T>>, T&>()
157 return *
reinterpret_cast<T*
>(data_);
159 return *
reinterpret_cast<T*
>(data_);
162 operator const T&()
const
167 return *
reinterpret_cast<const T*
>(data_);
169 return *
reinterpret_cast<const T*
>(data_);
175 return const_cast<T&
>(
static_cast<const NDArray&
>(*this).
scalar());
183 return *
reinterpret_cast<const T*
>(data_);
185 return *
reinterpret_cast<const T*
>(data_);
189 template <
typename... Idxs>
192 return const_cast<T&
>(
static_cast<const NDArray&
>(*this).
at(idxs...));
195 template <
typename... Idxs>
196 const T&
at(Idxs... idxs)
const
198 std::array<std::int64_t, MAX_DIMS> indices = {idxs...};
199 if (
sizeof...(Idxs) !=
static_cast<std::size_t
>(dims_))
202 return *
reinterpret_cast<const T*
>(data_);
204 for (std::size_t i = 0; i < indices.size(); ++i)
206 if (indices[i] < 0 || indices[i] >= shape_[i])
209 return *
reinterpret_cast<const T*
>(data_);
221 std::int64_t elemSize_;
222 std::int64_t numElements_;
224 void calculateSizeAndStrides()
228 for (std::int64_t i = 0; i < dims_; ++i)
232 numElements_ = size_ / elemSize_;
235 std::int64_t stride = 1;
236 for (std::int64_t i = dims_ - 1; i >= 0; --i)
238 strides_[i] = stride;
243 std::int64_t computeOffset(
const std::array<std::int64_t, MAX_DIMS>& idxs)
const
245 std::int64_t offset = 0;
246 for (std::int64_t i = 0; i < dims_; ++i)
248 offset += idxs[i] * strides_[i];
NDArray is a class that can be used in both host and device code. NDArray is automatically divided in...
const T & operator()(Idxs... idxs) const
std::int64_t elemSize() const
NDArray(T *data, std::initializer_list< int64_t > shape)
const StridesArray & strides() const
NDArray(void *data, const ShapeArray &shape, std::int64_t elemSize)
T & operator()(Idxs... idxs)
std::array< std::int64_t, MAX_DIMS > ShapeArray
std::int64_t numElements() const
const ShapeArray & shape() const
const T & at(Idxs... idxs) const
std::int64_t dims() const
NDArray< const T > operator[](std::int64_t idx) const
NDArray< T > operator[](std::int64_t idx)
std::int64_t size() const
NDArray(T *data, const ShapeArray &shape)
std::array< std::int64_t, MAX_DIMS > StridesArray
static constexpr std::int64_t MAX_DIMS