8 #include <initializer_list>
12 #include <type_traits>
34 : data_(const_cast<void*>(static_cast<const void*>(
data))),
35 dims_(static_cast<std::int64_t>(
shape.
size())),
38 std::fill(shape_.begin(), shape_.end(), 0);
39 std::fill(strides_.begin(), strides_.end(), 0);
40 for (
size_t i = 0; i <
shape.size(); ++i)
42 shape_[i] =
static_cast<std::int64_t
>(*(
shape.begin() + i));
44 calculateSizeAndStrides();
47 : data_(const_cast<void*>(reinterpret_cast<const void*>(
data))),
51 std::fill(shape_.begin(), shape_.end(), 0);
52 std::fill(strides_.begin(), strides_.end(), 0);
53 for (
size_t i = 0; i <
MAX_DIMS; ++i)
61 calculateSizeAndStrides();
69 std::fill(shape_.begin(), shape_.end(), 0);
70 std::fill(strides_.begin(), strides_.end(), 0);
71 for (
size_t i = 0; i <
MAX_DIMS; ++i)
79 calculateSizeAndStrides();
86 return reinterpret_cast<T*
>(data_);
90 return reinterpret_cast<const T*
>(data_);
123 template <
typename... Idxs>
126 return const_cast<T&
>(
static_cast<const NDArray&
>(*this)(idxs...));
129 template <
typename... Idxs>
134 std::vector<std::int64_t> indices = {idxs...};
135 if (indices.size() !=
static_cast<size_t>(dims_))
137 throw std::out_of_range(
"Number of indices does not match ndarray dimensions");
139 for (
size_t i = 0; i < indices.size(); ++i)
141 if (indices[i] < 0 || indices[i] >= shape_[i])
143 throw std::out_of_range(
"Index out of bounds");
147 return reinterpret_cast<const T*
>(data_)[computeOffset({idxs...})];
153 auto const_result =
static_cast<const NDArray&
>(*this)[idx];
155 const_cast<T*
>(const_result.data()),
156 const_result.shape());
163 if (idx < 0 || idx >= shape_[0])
165 throw std::out_of_range(
"Index out of bounds for first dimension");
169 std::fill(new_shape.begin(), new_shape.end(), 0);
170 std::copy(shape_.begin() + 1, shape_.end(), new_shape.begin());
173 reinterpret_cast<const T*
>(data_) + idx * strides_[0],
178 template <
typename U = T>
179 operator typename std::enable_if_t<std::is_same_v<U, std::remove_const_t<T>>, T&>()
183 throw std::logic_error(
"Scalar conversion only available for 0-dimensional ndarrays");
185 return *
reinterpret_cast<T*
>(data_);
188 operator const T&()
const
192 throw std::logic_error(
"Scalar conversion only available for 0-dimensional ndarrays");
194 return *
reinterpret_cast<const T*
>(data_);
200 return const_cast<T&
>(
static_cast<const NDArray&
>(*this).
scalar());
207 throw std::logic_error(
"scalar() only available for 0-dimensional ndarrays");
209 return *
reinterpret_cast<const T*
>(data_);
213 template <
typename... Idxs>
216 return const_cast<T&
>(
static_cast<const NDArray&
>(*this).
at(idxs...));
219 template <
typename... Idxs>
220 const T&
at(Idxs... idxs)
const
222 std::vector<std::int64_t> indices = {idxs...};
223 if (indices.size() !=
static_cast<size_t>(dims_))
225 throw std::out_of_range(
"Number of indices does not match ndarray dimensions");
227 for (
size_t i = 0; i < indices.size(); ++i)
229 if (indices[i] < 0 || indices[i] >= shape_[i])
231 throw std::out_of_range(
"Index out of bounds");
243 std::int64_t elemSize_;
244 std::int64_t numElements_;
246 void calculateSizeAndStrides()
250 for (std::int64_t i = 0; i < dims_; ++i)
254 numElements_ = size_ / elemSize_;
257 std::int64_t stride = 1;
258 for (std::int64_t i = dims_ - 1; i >= 0; --i)
260 strides_[i] = stride;
265 std::int64_t computeOffset(
const std::array<std::int64_t, MAX_DIMS>& idxs)
const
267 std::int64_t offset = 0;
268 for (std::int64_t i = 0; i < dims_; ++i)
270 offset += idxs[i] * strides_[i];
NDArray is a class that can be used in both host and device code. NDArray is automatically divided in...
static constexpr std::int64_t MAX_DIMS
NDArray< const T > operator[](std::int64_t idx) const
const T & at(Idxs... idxs) const
const T & operator()(Idxs... idxs) const
std::array< std::int64_t, MAX_DIMS > ShapeArray
NDArray< T > operator[](std::int64_t idx)
const StridesArray & strides() const
std::int64_t numElements() const
NDArray(void *data, const ShapeArray &shape, std::int64_t elemSize)
T & operator()(Idxs... idxs)
std::int64_t dims() const
NDArray(T *data, std::initializer_list< U > shape)
NDArray(T *data, const ShapeArray &shape)
std::int64_t size() const
std::array< std::int64_t, MAX_DIMS > StridesArray
std::int64_t elemSize() const
const ShapeArray & shape() const