8#include <initializer_list>
36 Shape(std::initializer_list<int64_t> dims) : dims_(dims) {
37 for (int64_t d : dims_) {
39 throw std::invalid_argument(
"Shape: dimensions must be non-negative");
53 explicit Shape(
const std::vector<int64_t>& dims) : dims_(dims) {
54 for (int64_t d : dims_) {
56 throw std::invalid_argument(
"Shape: dimensions must be non-negative");
62 [[nodiscard]] int64_t
rank()
const {
return static_cast<int64_t
>(dims_.size()); }
74 [[nodiscard]] int64_t
numel()
const {
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");
93 [[nodiscard]] int64_t
dim(
size_t index)
const {
106 [[nodiscard]]
bool operator==(
const Shape& other)
const {
return dims_ == other.dims_; }
110 std::vector<int64_t> dims_;
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