Stratax 0.2.0
Loading...
Searching...
No Matches
Shape.hpp
1#pragma once
2
3#include <vector>
4
5#include "Buffer.hpp"
6#include "../Concepts.hpp"
7#include "../Exceptions.hpp"
9
10#include <ostream>
11
12namespace stratax::core {
13
21class Shape
22{
23private:
25
26 void validate_dimensions() const
27 {
28 for (std::size_t dim : dims_)
29 {
30 validation::nonnegative_shape_dimension(
31 dim,
32 "Shape dimensions cannot be negative.");
33 }
34 }
35
36public:
38 struct allow_zero_t {};
39
41 static constexpr allow_zero_t allow_zero{};
42
45
48
51
54
60 Shape() noexcept = default;
61
71 template<Integral... Dims>
72 requires (sizeof...(Dims) > 0)
73 Shape(Dims... dims)
74 : dims_{static_cast<std::size_t>(dims)...}
75 {
76 validate_dimensions();
77 }
78
88 Shape(std::initializer_list<std::size_t> list, allow_zero_t allow_zero) : dims_(list)
89 {
90 (void)allow_zero;
91 }
92
99 : dims_(dims)
100 {
101 validate_dimensions();
102 }
103
111 : dims_(dims)
112 {
113 (void)allow_zero;
114 }
115
122 : dims_(std::move(dims))
123 {
124 validate_dimensions();
125 }
126
132 Shape(const std::vector<std::size_t>& dims)
133 : dims_(dims.size())
134 {
135 for (std::size_t i = 0; i < dims.size(); ++i)
136 {
137 dims_[i] = dims[i];
138 }
139 }
140
144 ~Shape() = default;
145
156 [[nodiscard]]
157 std::size_t elements() const
158 {
159 if (empty())
160 {
161 return 0;
162 }
163 std::size_t prod = 1;
164 for (std::size_t dim : dims_)
165 {
166 prod = validation::checked_multiply(prod, dim, "Shape elements overflow");
167 }
168 return prod;
169 }
170
176 [[nodiscard]] std::size_t rank() const
177 {
178 return dims_.size();
179 }
180
190 const std::size_t& operator()(std::size_t index) const
191 {
192 validation::require_index(index, rank(), "Shape dimension index out of bounds");
193 return dims_[index];
194 }
195
207 const std::size_t& operator[](std::ptrdiff_t index) const
208 {
209 return dims_[validation::normalize_index(index, rank(), "Shape dimension index out of bounds")];
210 }
211
217 [[nodiscard]] bool empty() const noexcept
218 {
219 return dims_.empty();
220 }
221
229 [[nodiscard]] bool operator==(const Shape& other) const noexcept
230 {
231 if (rank() != other.rank())
232 {
233 return false;
234 }
235 for (std::size_t i = 0; i < rank(); ++i)
236 {
237 if (dims_[i] != other.dims_[i])
238 {
239 return false;
240 }
241 }
242 return true;
243 }
244
252 [[nodiscard]] bool operator!=(const Shape& other) const noexcept
253 {
254 return !(*this == other);
255 }
256
262 [[nodiscard]] iterator begin() noexcept
263 {
264 return dims_.begin();
265 }
266
272 [[nodiscard]] iterator end() noexcept
273 {
274 return dims_.end();
275 }
276
282 [[nodiscard]] const_iterator begin() const noexcept
283 {
284 return dims_.begin();
285 }
286
292 [[nodiscard]] const_iterator end() const noexcept
293 {
294 return dims_.end();
295 }
296
302 [[nodiscard]] const_iterator cbegin() const noexcept
303 {
304 return dims_.cbegin();
305 }
306
312 [[nodiscard]] const_iterator cend() const noexcept
313 {
314 return dims_.cend();
315 }
316
322 [[nodiscard]] reverse_iterator rbegin() noexcept
323 {
324 return dims_.rbegin();
325 }
326
332 [[nodiscard]] const_reverse_iterator rbegin() const noexcept
333 {
334 return dims_.rbegin();
335 }
336
342 [[nodiscard]] const_reverse_iterator crbegin() const noexcept
343 {
344 return dims_.crbegin();
345 }
346
352 [[nodiscard]] reverse_iterator rend() noexcept
353 {
354 return dims_.rend();
355 }
356
362 [[nodiscard]] const_reverse_iterator rend() const noexcept
363 {
364 return dims_.rend();
365 }
366
372 [[nodiscard]] const_reverse_iterator crend() const noexcept
373 {
374 return dims_.crend();
375 }
376
382 void swap(Shape& other) noexcept
383 {
384 dims_.swap(other.dims_);
385 }
386
387};
388
397inline std::ostream& operator<<(std::ostream& os, const Shape& shape)
398{
399 os << "(";
400
401 bool first = true;
402 for (std::size_t dim : shape)
403 {
404 if (!first)
405 os << ", ";
406
407 os << dim;
408 first = false;
409 }
410
411 if (shape.rank() == 1)
412 {
413 os << ",";
414 }
415
416 os << ")";
417
418 return os;
419}
420
421}
422
Shared runtime validation helpers.
Owns contiguous dynamically allocated storage.
Definition Buffer.hpp:38
std::reverse_iterator< iterator > reverse_iterator
Mutable reverse iterator over contiguous buffer elements.
Definition Buffer.hpp:53
const T * const_iterator
Const iterator over contiguous buffer elements.
Definition Buffer.hpp:50
T * iterator
Mutable iterator over contiguous buffer elements.
Definition Buffer.hpp:47
std::reverse_iterator< const_iterator > const_reverse_iterator
Const reverse iterator over contiguous buffer elements.
Definition Buffer.hpp:56
~Shape()=default
Destroys the shape.
Shape(const std::vector< std::size_t > &dims)
Creates a shape by copying dimension lengths from a standard vector.
Definition Shape.hpp:132
std::size_t elements() const
Returns the total number of elements described by the shape.
Definition Shape.hpp:157
const std::size_t & operator[](std::ptrdiff_t index) const
Returns the length of a specific dimension using signed indexing.
Definition Shape.hpp:207
Buffer< std::size_t >::const_iterator const_iterator
Const iterator over dimension lengths.
Definition Shape.hpp:47
bool operator!=(const Shape &other) const noexcept
Returns whether two shapes differ in rank or dimension values.
Definition Shape.hpp:252
reverse_iterator rbegin() noexcept
Returns a reverse iterator to the last stored dimension.
Definition Shape.hpp:322
Buffer< std::size_t >::reverse_iterator reverse_iterator
Mutable reverse iterator over dimension lengths.
Definition Shape.hpp:50
reverse_iterator rend() noexcept
Returns a reverse iterator before the first stored dimension.
Definition Shape.hpp:352
Buffer< std::size_t >::iterator iterator
Mutable iterator over dimension lengths.
Definition Shape.hpp:44
const_iterator cend() const noexcept
Returns a const iterator one past the last stored dimension.
Definition Shape.hpp:312
const std::size_t & operator()(std::size_t index) const
Returns the length of a specific dimension.
Definition Shape.hpp:190
bool empty() const noexcept
Returns whether the shape has no dimensions.
Definition Shape.hpp:217
bool operator==(const Shape &other) const noexcept
Compares two shapes for exact rank and dimension equality.
Definition Shape.hpp:229
void swap(Shape &other) noexcept
Swaps the stored dimensions with another shape.
Definition Shape.hpp:382
const_reverse_iterator crbegin() const noexcept
Returns a const reverse iterator to the last stored dimension.
Definition Shape.hpp:342
const_iterator begin() const noexcept
Returns a const iterator to the first stored dimension.
Definition Shape.hpp:282
const_iterator end() const noexcept
Returns a const iterator one past the last stored dimension.
Definition Shape.hpp:292
Shape() noexcept=default
Creates an empty shape.
std::size_t rank() const
Returns the number of stored dimensions.
Definition Shape.hpp:176
iterator end() noexcept
Returns an iterator one past the last stored dimension.
Definition Shape.hpp:272
Shape(const Buffer< std::size_t > &dims)
Creates a shape by copying dimension lengths from a buffer.
Definition Shape.hpp:98
const_reverse_iterator crend() const noexcept
Returns a const reverse iterator before the first stored dimension.
Definition Shape.hpp:372
static constexpr allow_zero_t allow_zero
Tag value documenting that zero-valued dimensions are intentional.
Definition Shape.hpp:41
Buffer< std::size_t >::const_reverse_iterator const_reverse_iterator
Const reverse iterator over dimension lengths.
Definition Shape.hpp:53
const_iterator cbegin() const noexcept
Returns a const iterator to the first stored dimension.
Definition Shape.hpp:302
const_reverse_iterator rbegin() const noexcept
Returns a const reverse iterator to the last stored dimension.
Definition Shape.hpp:332
Shape(Buffer< std::size_t > &&dims)
Creates a shape by taking ownership of dimension lengths from a buffer.
Definition Shape.hpp:121
const_reverse_iterator rend() const noexcept
Returns a const reverse iterator before the first stored dimension.
Definition Shape.hpp:362
Shape(std::initializer_list< std::size_t > list, allow_zero_t allow_zero)
Creates a shape from dimension lengths using the explicit zero-dimension tag.
Definition Shape.hpp:88
iterator begin() noexcept
Returns an iterator to the first stored dimension.
Definition Shape.hpp:262
Shape(const Buffer< std::size_t > &dims, allow_zero_t allow_zero)
Creates a shape by copying dimension lengths using the explicit zero-dimension tag.
Definition Shape.hpp:110
Matches signed and unsigned integral types, excluding character-like types.
Definition Concepts.hpp:55
Tag type documenting that zero-valued dimensions are intentional.
Definition Shape.hpp:38