pulsatrix
Loading...
Searching...
No Matches
shape.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <cstdint>
8#include <initializer_list>
9#include <limits>
10#include <numeric>
11#include <stdexcept>
12#include <vector>
13
14#include "pulsatrix/assert.hpp"
15
16namespace pulsatrix {
17
24class Shape {
25public:
36 Shape(std::initializer_list<int64_t> dims) : dims_(dims) {
37 for (int64_t d : dims_) {
38 if (d < 0) {
39 throw std::invalid_argument("Shape: dimensions must be non-negative");
40 }
41 }
42 }
43
53 explicit Shape(const std::vector<int64_t>& dims) : dims_(dims) {
54 for (int64_t d : dims_) {
55 if (d < 0) {
56 throw std::invalid_argument("Shape: dimensions must be non-negative");
57 }
58 }
59 }
60
62 [[nodiscard]] int64_t rank() const { return static_cast<int64_t>(dims_.size()); }
63
74 [[nodiscard]] int64_t numel() const {
75 int64_t result = 1;
76 for (int64_t d : dims_) {
77 if (d != 0 && result > std::numeric_limits<int64_t>::max() / (d == 0 ? 1 : d)) {
78 throw std::overflow_error("Shape::numel: dimension product overflows int64_t");
79 }
80 result *= d;
81 }
82 return result;
83 }
84
93 [[nodiscard]] int64_t dim(size_t index) const {
94 PULSATRIX_ASSERT(index < dims_.size());
95 return dims_[index];
96 }
97
104 [[nodiscard]] bool is_reshape_compatible(const Shape& other) const { return numel() == other.numel(); }
105
106 [[nodiscard]] bool operator==(const Shape& other) const { return dims_ == other.dims_; }
107 [[nodiscard]] bool operator!=(const Shape& other) const { return !(*this == other); }
108
109private:
110 std::vector<int64_t> dims_;
111};
112
113} // namespace pulsatrix
PULSATRIX_ASSERT – debug-only invariant check for programmer errors, distinct from throw (used for ca...
#define PULSATRIX_ASSERT(cond)
Aborts with a diagnostic message if cond is false. Debug-only – use for conditions that indicate a bu...
Definition assert.hpp:22
An N-dimensional shape. A plain aggregate of dimensions with no invariant beyond "non-negative dimens...
Definition shape.hpp:24
int64_t dim(size_t index) const
Size of a single dimension.
Definition shape.hpp:93
int64_t numel() const
Total element count – the product of all dimensions.
Definition shape.hpp:74
Shape(std::initializer_list< int64_t > dims)
Constructs a shape from a dimension list. An empty list is a rank-0 scalar (numel() == 1).
Definition shape.hpp:36
bool operator==(const Shape &other) const
Definition shape.hpp:106
Shape(const std::vector< int64_t > &dims)
Constructs a shape from a runtime-sized dimension list.
Definition shape.hpp:53
bool operator!=(const Shape &other) const
Definition shape.hpp:107
bool is_reshape_compatible(const Shape &other) const
Whether this shape and other have the same numel() – the precondition for a valid reshape between the...
Definition shape.hpp:104
int64_t rank() const
Number of dimensions. 0 for a scalar.
Definition shape.hpp:62
Definition acquisition_functions.hpp:16