PXL
ndarray.hpp
Go to the documentation of this file.
1 // SPDX-License-Identifier: Apache-2.0
2 // Copyright 2024 XCENA Inc.
3 
4 #pragma once
5 
6 #include <array>
7 #include <cstdint>
8 #include <initializer_list>
9 #include <numeric>
10 #include <stdexcept>
11 #include <string>
12 #include <type_traits>
13 #include <vector>
14 
15 namespace pxl
16 {
17 
23 template <typename T>
24 class NDArray
25 {
26 public:
27  static constexpr std::int64_t MAX_DIMS = 10;
28  using ShapeArray = std::array<std::int64_t, MAX_DIMS>;
29  using StridesArray = std::array<std::int64_t, MAX_DIMS>;
30 
31  NDArray() = default;
32  template <typename U>
33  NDArray(T* data, std::initializer_list<U> shape)
34  : data_(const_cast<void*>(static_cast<const void*>(data))),
35  dims_(static_cast<std::int64_t>(shape.size())),
36  elemSize_(sizeof(T))
37  {
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)
41  {
42  shape_[i] = static_cast<std::int64_t>(*(shape.begin() + i));
43  }
44  calculateSizeAndStrides();
45  }
47  : data_(const_cast<void*>(reinterpret_cast<const void*>(data))),
48  dims_(0),
49  elemSize_(sizeof(T))
50  {
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)
54  {
55  if (shape[i] != 0)
56  {
57  shape_[i] = shape[i];
58  dims_++;
59  }
60  }
61  calculateSizeAndStrides();
62  }
63 
64  NDArray(void* data, const ShapeArray& shape, std::int64_t elemSize)
65  : data_(data),
66  dims_(0),
67  elemSize_(elemSize)
68  {
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)
72  {
73  if (shape[i] != 0)
74  {
75  shape_[i] = shape[i];
76  dims_++;
77  }
78  }
79  calculateSizeAndStrides();
80  }
81 
82  ~NDArray() = default;
83 
84  T* data()
85  {
86  return reinterpret_cast<T*>(data_);
87  }
88  const T* data() const
89  {
90  return reinterpret_cast<const T*>(data_);
91  }
92 
93  void setData(void* data)
94  {
95  data_ = data;
96  }
97 
98  const ShapeArray& shape() const
99  {
100  return shape_;
101  }
102  const StridesArray& strides() const
103  {
104  return strides_;
105  }
106  std::int64_t dims() const
107  {
108  return dims_;
109  }
110  std::int64_t elemSize() const
111  {
112  return elemSize_;
113  }
114  std::int64_t size() const
115  {
116  return size_;
117  }
118  std::int64_t numElements() const
119  {
120  return numElements_;
121  }
122 
123  template <typename... Idxs>
124  T& operator()(Idxs... idxs)
125  {
126  return const_cast<T&>(static_cast<const NDArray&>(*this)(idxs...));
127  }
128 
129  template <typename... Idxs>
130  const T& operator()(Idxs... idxs) const
131  {
132 #ifndef NDEBUG
133  // Debug mode: perform bounds checking
134  std::vector<std::int64_t> indices = {idxs...};
135  if (indices.size() != static_cast<size_t>(dims_))
136  {
137  throw std::out_of_range("Number of indices does not match ndarray dimensions");
138  }
139  for (size_t i = 0; i < indices.size(); ++i)
140  {
141  if (indices[i] < 0 || indices[i] >= shape_[i])
142  {
143  throw std::out_of_range("Index out of bounds");
144  }
145  }
146 #endif
147  return reinterpret_cast<const T*>(data_)[computeOffset({idxs...})];
148  }
149 
150  // Subndarray access with operator[] - always returns a NDArray
151  NDArray<T> operator[](std::int64_t idx)
152  {
153  auto const_result = static_cast<const NDArray&>(*this)[idx];
154  return NDArray<T>(
155  const_cast<T*>(const_result.data()),
156  const_result.shape());
157  }
158 
159  NDArray<const T> operator[](std::int64_t idx) const
160  {
161 #ifndef NDEBUG
162  // Debug mode: perform bounds checking
163  if (idx < 0 || idx >= shape_[0])
164  {
165  throw std::out_of_range("Index out of bounds for first dimension");
166  }
167 #endif
168  ShapeArray new_shape;
169  std::fill(new_shape.begin(), new_shape.end(), 0);
170  std::copy(shape_.begin() + 1, shape_.end(), new_shape.begin());
171 
172  return NDArray<const T>(
173  reinterpret_cast<const T*>(data_) + idx * strides_[0],
174  new_shape);
175  }
176 
177  // Implicit conversion to scalar for 0-dimensional ndarrays
178  template <typename U = T>
179  operator typename std::enable_if_t<std::is_same_v<U, std::remove_const_t<T>>, T&>()
180  {
181  if (dims_ != 0)
182  {
183  throw std::logic_error("Scalar conversion only available for 0-dimensional ndarrays");
184  }
185  return *reinterpret_cast<T*>(data_);
186  }
187 
188  operator const T&() const
189  {
190  if (dims_ != 0)
191  {
192  throw std::logic_error("Scalar conversion only available for 0-dimensional ndarrays");
193  }
194  return *reinterpret_cast<const T*>(data_);
195  }
196 
197  // Explicit scalar() method for clearer intent
198  T& scalar()
199  {
200  return const_cast<T&>(static_cast<const NDArray&>(*this).scalar());
201  }
202 
203  const T& scalar() const
204  {
205  if (dims_ != 0)
206  {
207  throw std::logic_error("scalar() only available for 0-dimensional ndarrays");
208  }
209  return *reinterpret_cast<const T*>(data_);
210  }
211 
212  // Multi-dimensional access with bounds checking
213  template <typename... Idxs>
214  T& at(Idxs... idxs)
215  {
216  return const_cast<T&>(static_cast<const NDArray&>(*this).at(idxs...));
217  }
218 
219  template <typename... Idxs>
220  const T& at(Idxs... idxs) const
221  {
222  std::vector<std::int64_t> indices = {idxs...};
223  if (indices.size() != static_cast<size_t>(dims_))
224  {
225  throw std::out_of_range("Number of indices does not match ndarray dimensions");
226  }
227  for (size_t i = 0; i < indices.size(); ++i)
228  {
229  if (indices[i] < 0 || indices[i] >= shape_[i])
230  {
231  throw std::out_of_range("Index out of bounds");
232  }
233  }
234  return operator()(idxs...);
235  }
236 
237 private:
238  void* data_;
239  ShapeArray shape_;
240  StridesArray strides_;
241  std::int64_t dims_;
242  std::int64_t size_;
243  std::int64_t elemSize_;
244  std::int64_t numElements_;
245 
246  void calculateSizeAndStrides()
247  {
248  // Calculate size
249  size_ = elemSize_;
250  for (std::int64_t i = 0; i < dims_; ++i)
251  {
252  size_ *= shape_[i];
253  }
254  numElements_ = size_ / elemSize_;
255 
256  // Calculate strides
257  std::int64_t stride = 1;
258  for (std::int64_t i = dims_ - 1; i >= 0; --i)
259  {
260  strides_[i] = stride;
261  stride *= shape_[i];
262  }
263  }
264 
265  std::int64_t computeOffset(const std::array<std::int64_t, MAX_DIMS>& idxs) const
266  {
267  std::int64_t offset = 0;
268  for (std::int64_t i = 0; i < dims_; ++i)
269  {
270  offset += idxs[i] * strides_[i];
271  }
272  return offset;
273  }
274 };
275 
276 } // namespace pxl
NDArray is a class that can be used in both host and device code. NDArray is automatically divided in...
Definition: ndarray.hpp:25
static constexpr std::int64_t MAX_DIMS
Definition: ndarray.hpp:27
NDArray< const T > operator[](std::int64_t idx) const
Definition: ndarray.hpp:159
const T * data() const
Definition: ndarray.hpp:88
const T & at(Idxs... idxs) const
Definition: ndarray.hpp:220
const T & scalar() const
Definition: ndarray.hpp:203
const T & operator()(Idxs... idxs) const
Definition: ndarray.hpp:130
std::array< std::int64_t, MAX_DIMS > ShapeArray
Definition: ndarray.hpp:28
NDArray< T > operator[](std::int64_t idx)
Definition: ndarray.hpp:151
T * data()
Definition: ndarray.hpp:84
const StridesArray & strides() const
Definition: ndarray.hpp:102
T & at(Idxs... idxs)
Definition: ndarray.hpp:214
std::int64_t numElements() const
Definition: ndarray.hpp:118
NDArray(void *data, const ShapeArray &shape, std::int64_t elemSize)
Definition: ndarray.hpp:64
T & operator()(Idxs... idxs)
Definition: ndarray.hpp:124
std::int64_t dims() const
Definition: ndarray.hpp:106
T & scalar()
Definition: ndarray.hpp:198
NDArray(T *data, std::initializer_list< U > shape)
Definition: ndarray.hpp:33
NDArray(T *data, const ShapeArray &shape)
Definition: ndarray.hpp:46
std::int64_t size() const
Definition: ndarray.hpp:114
std::array< std::int64_t, MAX_DIMS > StridesArray
Definition: ndarray.hpp:29
std::int64_t elemSize() const
Definition: ndarray.hpp:110
const ShapeArray & shape() const
Definition: ndarray.hpp:98
NDArray()=default
~NDArray()=default
void setData(void *data)
Definition: ndarray.hpp:93
Definition: config.hpp:11