Source
stdlib/ndarray/ndarray.hpp
1
// Copyright (c) 2026 BigBrain LLC. MIT-licensed (see LICENSE).2
// Original work; see ACKNOWLEDGMENTS.md for the open-source ideas we build upon.3
#pragma once5
/**6
* @file ndarray.hpp7
* @brief cheatah `ndarray` — our own numpy-flavored N-dimensional array8
* (`basic_ndarray<T>` over any @ref Element type; `NDArray` is the `double`9
* default) with NumPy broadcasting, surfaced as a `NDArray` class plus free10
* functions (a .purr program writes `ndarray.zeros([2, 3])`).11
* See https://numpy.org/doc/stable/user/basics.broadcasting.html.12
*13
* `import ndarray` includes this header and links `libcheatah_ndarray`. Unit tests:14
* `stdlib/tests/ndarray_test.cpp`; the suite runs under AddressSanitizer (the `asan`15
* preset) and Valgrind (`security/run-valgrind.sh`) on every QA-gate run.16
*17
* @note Design (the "pointers + a bit of thinking"): the elements live in a shared18
* buffer (`std::shared_ptr<buffer_t<T>>`) and an array is a VIEW into19
* it — {shape, strides, offset}. That makes reshape and **broadcast**20
* zero-copy: to stretch a dimension of size 1 we give it a stride of 0, so21
* every index along it reads the same element. Shared ownership = memory-safe,22
* no manual frees. `size` below is the element count (product of dims).23
*/24
#include <algorithm>25
#include <array>26
#include <cmath>27
#include <complex>28
#include <concepts>29
#include <cstddef>30
#include <initializer_list>31
#include <limits>32
#include <memory>33
#include <new> // placement new, for the default-init buffer allocator34
#include <numeric>35
#include <sstream>36
#include <stdexcept>37
#include <string>38
#include <type_traits>39
#include <utility> // std::forward / std::move40
#include <version> // __cpp_lib_execution feature-test macro41
#include <vector>43
// The unsequenced execution policy lets std::transform/std::reduce vectorize. libstdc++44
// provides it; Apple's libc++ historically ships no usable <execution>, so guard on the45
// feature-test macro and fall back to the plain (policy-less) overloads where it's46
// absent. This is speed-neutral: `unseq` is unsequenced (no threads, no TBB) and for47
// these simple element-wise loops the -O3 -march=native auto-vectorizer produces the48
// same SIMD either way — the transcendental vectorization comes from ufunc_simd.cpp's49
// libmvec/Accelerate kernels, not from this policy.50
#if defined(__cpp_lib_execution)51
#include <execution>52
#define CHEATAH_UNSEQ std::execution::unseq,53
#else54
#define CHEATAH_UNSEQ55
#endif57
namespace cheatah::ndarray {59
/// Numeric<T>: an arithmetic element type an ndarray can store (int or float60
/// family). Storage, construction, and elementwise +-* require only this.61
template <typename T>62
concept Numeric = std::is_arithmetic_v<T>;63
/// FloatingPoint<T>: a real floating type. The linalg decompositions (solve, inv,64
/// det, svd, eig) need division/√, so they constrain to this — calling them on an65
/// integer array fails with a clear "FloatingPoint not satisfied", not template spam.66
template <typename T>67
concept FloatingPoint = std::floating_point<T>;69
/// @cond INTERNAL70
template <typename T>71
struct is_complex : std::false_type {};72
template <typename U>73
struct is_complex<std::complex<U>> : std::bool_constant<std::is_floating_point_v<U>> {};74
/// @endcond76
/// Whether `T` is a `std::complex` of a floating type — the trait behind @ref Field.77
template <typename T>78
inline constexpr bool is_complex_v = is_complex<T>::value;80
/// Field<T>: a scalar an ndarray can store — a real arithmetic type OR a81
/// `std::complex` of a floating type. This is what makes **complex** matrices and82
/// vectors first-class (Hermitian operators, complex wavefunctions), and lets a83
/// REAL matrix yield the COMPLEX eigenvalues it mathematically has.84
template <typename T>85
concept Field = std::is_arithmetic_v<T> || is_complex_v<T>;87
/// Element<T>: the broadest bound — any type an ndarray may STORE. A @ref Field (real/complex88
/// number) OR any MOVABLE type, so a fixed-size struct (a 2-D point, an RGBA colour, a GPU89
/// vertex) lives in an ndarray too. Elements are MOVED into the buffer on construction and the90
/// backing buffer is never silently deep-copied (copying an ndarray shares the buffer — an O(1)91
/// view). The arithmetic surface (elementwise ops, ufuncs, reductions, linalg) stays constrained92
/// to @ref Field and the duplicating factories to @ref Copyable, so a move-only element still93
/// stores / indexes / views / moves — it simply cannot be summed or deep-copied, and the compiler94
/// says so by design (cheatah discourages hidden copies on hot data).95
template <typename T>96
concept Element = Field<T> || std::movable<T>;98
/// Copyable<T>: an @ref Element that may ALSO be duplicated. It gates only the value-fill99
/// factories (full / full_like) and reshape's deep copy — the paths that replicate elements.100
/// Numbers and ordinary copyable POD structs satisfy it; a move-only struct does NOT, so an101
/// accidental deep copy of it fails to compile. Copying an ndarray CONTAINER never needs this —102
/// it is always a shared-buffer view (and the GPU borrows that buffer in place, never copying it).103
template <typename T>104
concept Copyable = Element<T> && std::copyable<T>;106
/// Subscript<T>: what may address an axis — an integer, OR a scoped `enum class` whose ordinal names107
/// the position (a column label). This concept is the ONLY door through which a scoped enum becomes an108
/// integer: `enum class` values stay strongly typed everywhere else, and the implicit109
/// enum-to-index conversion is confined to array subscripting, exactly where a named column belongs.110
template <typename T>111
concept Subscript = std::is_convertible_v<T, long long> || std::is_enum_v<T>;113
/// The integer position an @ref Subscript addresses. For a scoped enum this is its underlying ordinal;114
/// this `static_cast` is the whole of the enum-to-index conversion the language sanctions.115
/// @tparam Ix the subscript type: an integer, or a scoped `enum class`.116
/// @param i the subscript to resolve.117
/// @return the integer position it names (a scoped enum's underlying ordinal).118
/// @complexity O(1). @alloc none.119
/// @test Fixarray.EnumIndexingOnVectorsAndMatrices120
template <Subscript Ix>121
[[nodiscard]] constexpr long long subscript_index(Ix i) noexcept {122
return static_cast<long long>(i);123
}125
/// @cond INTERNAL126
template <typename T>127
struct real_base {128
using type = T;129
};130
template <typename U>131
struct real_base<std::complex<U>> {132
using type = U;133
};134
/// @endcond136
/// The real type underlying a @ref Field `T` (`double` for both `double` and137
/// `complex<double>`).138
template <typename T>139
using real_base_t = typename real_base<T>::type;141
/// complex_of_t<T>: the complex type over T's real base. `eig`/`eigvals` return an142
/// array of these, because a real matrix can have complex eigenvalues (conjugate143
/// pairs) — e.g. the rotation matrix [[0,-1],[1,0]] has eigenvalues ±i.144
template <typename T>145
using complex_of_t = std::complex<real_base_t<T>>;147
namespace detail {148
/// @cond INTERNAL149
/// An allocator identical to `std::allocator<T>` in every respect EXCEPT that150
/// DEFAULT (no-value) construction — what `vector(n)` / `resize(n)` perform — leaves a151
/// trivially-constructible element UNINITIALIZED instead of value-initializing it to 0.152
///153
/// Every ndarray op that allocates a result buffer it then overwrites in full (binary154
/// ops, ufuncs, reshape, array()) would otherwise pay a throwaway zero-fill of the whole155
/// buffer first. That wasted write pass is hidden on compute-heavy ops but DOMINATES156
/// bandwidth-bound ones — `add` was ≈1.5× of NumPy purely from the extra memset. With157
/// this allocator the sizing path skips it; the value-filling forms (`assign(n, v)`,158
/// `vector(n, v)` used by zeros/full/scalar) are untouched and still initialize.159
template <typename T>160
struct default_init_allocator : std::allocator<T> {161
using std::allocator<T>::allocator;162
template <typename U>163
struct rebind {164
using other = default_init_allocator<U>;165
};166
/// Default construction: a trivially-constructible element is left uninitialized.167
template <typename U>168
void construct(U* p) noexcept(std::is_nothrow_default_constructible_v<U>) {169
::new (static_cast<void*>(p)) U; // default-init (no `()`): no zeroing for trivial U170
}171
/// Every other construction (value-fill, copy, emplace) behaves exactly as normal.172
template <typename U, typename... Args>173
void construct(U* p, Args&&... args) {174
::new (static_cast<void*>(p)) U(std::forward<Args>(args)...);175
}176
};177
/// @endcond179
/// C-order (row-major) strides for a shape.180
inline std::vector<std::ptrdiff_t> contiguous_strides(const std::vector<std::size_t>& shape) {181
std::vector<std::ptrdiff_t> s(shape.size());182
std::ptrdiff_t step = 1;183
for (std::size_t i = shape.size(); i-- > 0;) {184
s[i] = step;185
step *= static_cast<std::ptrdiff_t>(shape[i]);186
}187
return s;188
}189
/// Overflow-checked product of the dimensions (a wrapped size_t would under-allocate,190
/// turning later element access into out-of-bounds writes — reject it up front).191
inline std::size_t product(const std::vector<std::size_t>& shape) {192
std::size_t t = 1;193
for (std::size_t d : shape) {194
if (d != 0 && t > std::numeric_limits<std::size_t>::max() / d) {195
throw std::runtime_error("ndarray: shape too large (size overflow)");196
}197
t *= d;198
}199
return t;200
}201
/// Convert signed dims/indices to sizes, rejecting negatives (a negative cast to202
/// size_t becomes huge -> under-allocation / OOB). Validate at the boundary.203
inline std::vector<std::size_t> to_size(const std::vector<long long>& v) {204
std::vector<std::size_t> out(v.size());205
for (std::size_t i = 0; i < v.size(); ++i) {206
if (v[i] < 0) throw std::runtime_error("ndarray: negative dimension or index");207
out[i] = static_cast<std::size_t>(v[i]);208
}209
return out;210
}211
/// Advance a C-order multi-index odometer; false when it wraps past the end.212
inline bool next_index(std::vector<std::size_t>& idx, const std::vector<std::size_t>& shape) {213
for (std::size_t i = shape.size(); i-- > 0;) {214
if (++idx[i] < shape[i]) return true;215
idx[i] = 0;216
}217
return false;218
}220
/// Peel `std::vector<>` layers off a (possibly deeply nested) list type to reach the221
/// leaf scalar — `nested_scalar_t<std::vector<std::vector<double>>>` is `double`.222
template <typename V> struct nested_scalar { using type = V; };223
template <typename U> struct nested_scalar<std::vector<U>> {224
using type = typename nested_scalar<U>::type;225
};226
template <typename V> using nested_scalar_t = typename nested_scalar<V>::type;228
/// Whether `V` is a `std::vector<…>` (used to tell a nested list from a scalar leaf).229
template <typename V> inline constexpr bool is_std_vector_v = false;230
template <typename U> inline constexpr bool is_std_vector_v<std::vector<U>> = true;232
/// Flatten a scalar leaf into the C-order buffer (recursion base case).233
template <Element T>234
void nested_collect(T x, std::vector<T>& flat, std::vector<std::size_t>& /*unused*/, std::size_t /*unused*/) {235
flat.push_back(std::move(x));236
}237
/// Walk a nested list: record each axis length the first time it is seen, reject a238
/// ragged list (a row whose length differs from its siblings — numpy does too), and239
/// flatten the leaves in C-order. The leaf scalar must be a @ref Field.240
template <typename U>241
requires Field<nested_scalar_t<U>>242
void nested_collect(const std::vector<U>& v, std::vector<nested_scalar_t<U>>& flat,243
std::vector<std::size_t>& shape, std::size_t depth) {244
if (depth == shape.size()) shape.push_back(v.size());245
else if (shape[depth] != v.size())246
throw std::runtime_error("ndarray: array(...) ragged nested list (a row's length "247
"differs from its siblings)");248
for (const U& e : v) nested_collect(e, flat, shape, depth + 1);249
}250
} // namespace detail252
/// The backing store of an ndarray: a flat, contiguous, shared element buffer. It uses253
/// @ref detail::default_init_allocator so a freshly-sized result buffer that an op is254
/// about to overwrite in full is not needlessly zero-filled first. zeros/full/scalar,255
/// which value-fill, are unaffected — only the no-value sizing path skips initialization.256
template <Element T>257
using buffer_t = std::vector<T, detail::default_init_allocator<T>>;259
/**260
* @brief An N-dimensional array of `T` (a @ref Field element type — real or complex):261
* a view ({shape, strides, offset}) over a shared element buffer.262
*263
* Copies are cheap and share the buffer; reshape/broadcast produce new views without264
* copying elements. Index math goes through @ref at, which bounds-checks. The element265
* type is deduced from the data (e.g. `array([1,2,3])` is integer, `array([1.0,…])`266
* is double); `NDArray` is the default `basic_ndarray<double>`. Complex element types267
* (`std::complex<double>`) make complex matrices/vectors — and the complex eigenvalues268
* a real matrix can have — first-class.269
*/270
template <Element T>271
class basic_ndarray {272
public:273
using value_type = T; ///< The stored element type (an @ref Element `T`).274
/**275
* Construct an empty 0-d array with a fresh empty buffer.276
*277
* Leaves shape and strides empty and the buffer holding no elements; note this278
* is distinct from a 0-d scalar (see @ref scalar), whose buffer holds one element.279
* @complexity O(1).280
* @alloc allocates the empty shared buffer.281
* @test CheatahNDArray.ToStringScalar282
* @systest StdlibE2E.Ndarray283
*/284
basic_ndarray() : data_(std::make_shared<buffer_t<T>>()) {}285
/**286
* Construct a contiguous array of @p shape filled with @p fill.287
*288
* Allocates a fresh buffer of `product(shape)` elements all set to @p fill and289
* computes C-order (row-major) strides; the dimension product is overflow-checked.290
* @param shape the dimensions.291
* @param fill value for every element.292
* @complexity O(size).293
* @alloc allocates a new shared buffer (`shared_ptr<buffer_t<T>>`) of294
* `product(shape)` elements; throws if the shape overflows size_t.295
* @test CheatahNDArray.ShapeFactoriesAndReductions296
* @systest StdlibE2E.Ndarray297
*/298
explicit basic_ndarray(std::vector<std::size_t> shape, T fill = T{}) // contiguous299
: data_(std::make_shared<buffer_t<T>>()), shape_(std::move(shape)) {300
// resize (default-init: no zero pass) then std::fill — the fill goes through301
// operator= on built elements, which keeps libstdc++'s memset/SIMD fast path.302
// (vector::assign(n, v) would route through the allocator's construct, which a303
// default-init allocator forces element-by-element — measurably slower.)304
data_->resize(detail::product(shape_));305
std::fill(data_->begin(), data_->end(), fill);306
strides_ = detail::contiguous_strides(shape_);307
}308
/**309
* Build a contiguous array of @p shape whose buffer is sized but left310
* UNINITIALIZED — for internal ops (binary ops, ufuncs, reshape, array) that311
* immediately write every element, so paying for a zero-fill first is pure waste.312
* @param shape the dimensions.313
* @return an array of @p shape with an uninitialized buffer.314
* @complexity O(ndim) beyond the allocation (no element initialization pass for a trivially-constructible `T`).315
* @alloc allocates an uninitialized `product(shape)`-element buffer; overflow-checked.316
* @test CheatahNDArray.BroadcastingAdd317
* @systest StdlibE2E.Ndarray318
*/319
static basic_ndarray uninitialized(std::vector<std::size_t> shape) {320
basic_ndarray out;321
out.shape_ = std::move(shape);322
out.data_->resize(detail::product(out.shape_)); // default-init alloc -> no zero-fill323
out.strides_ = detail::contiguous_strides(out.shape_);324
return out;325
}326
/**327
* Construct a view from explicit buffer/shape/strides/offset (used by views).328
*329
* Stores the supplied members verbatim with no validation or copy, so the new330
* array shares ownership of @p data; callers (e.g. @ref broadcast_to) are331
* responsible for passing strides/offset that stay within the buffer.332
* @param data the shared element buffer.333
* @param shape the dimensions.334
* @param strides element strides per dimension.335
* @param offset starting flat offset into @p data.336
* @complexity O(1).337
* @alloc shares @p data, no element copy.338
* @test CheatahNDArray.BroadcastTo339
* @systest StdlibE2E.Ndarray340
*/341
basic_ndarray(std::shared_ptr<buffer_t<T>> data, std::vector<std::size_t> shape,342
std::vector<std::ptrdiff_t> strides, std::size_t offset)343
: data_(std::move(data)), shape_(std::move(shape)), strides_(std::move(strides)),344
offset_(offset) {}346
/**347
* The shape (dimensions).348
* @return reference to the shape vector.349
* @complexity O(1).350
* @alloc none.351
* @test CheatahNDArray.ShapeFactoriesAndReductions352
* @systest StdlibE2E.Ndarray353
*/354
const std::vector<std::size_t>& shape() const { return shape_; }355
/**356
* The element strides.357
* @return reference to the strides vector.358
* @complexity O(1).359
* @alloc none.360
* @test CheatahNDArray.BroadcastTo361
*/362
const std::vector<std::ptrdiff_t>& strides() const { return strides_; }363
/**364
* The number of dimensions (rank).365
* @return `shape().size()`.366
* @complexity O(1).367
* @alloc none.368
* @test CheatahNDArray.BroadcastingAdd369
*/370
std::size_t ndim() const { return shape_.size(); }371
/**372
* The element count (product of dims; 1 for a 0-d scalar).373
*374
* Recomputes the overflow-checked product of the shape on each call rather than375
* caching it; an empty (no-dimension) shape yields 1.376
* @return the number of elements.377
* @complexity O(ndim).378
* @alloc none; throws on size overflow.379
* @test CheatahNDArray.ShapeFactoriesAndReductions380
*/381
std::size_t size() const { return detail::product(shape_); } // 1 for a 0-d scalar383
/**384
* MUTABLE element reference by multi-index — the write path behind cheatah385
* subscript assignment `x[i] = v` / `x[i, j] = v`. Negative indices count386
* from the end of their dimension; rank and bounds are checked.387
* @param ixs one (possibly negative) coordinate per dimension.388
* @return a writable reference into the backing buffer.389
* @complexity O(ndim).390
* @alloc none.391
* @test CheatahNDArray.SubscriptReadWrite392
* @crtest LangFeatures.NdarraySubscript393
*/394
template <typename... Ix>395
requires((Subscript<Ix> && ...) && sizeof...(Ix) > 0)396
T& item_ref(Ix... ixs) {397
return (*data_)[item_pos(ixs...)];398
}400
/**401
* READ-ONLY element reference by multi-index — the same rank/bounds-checked position math as402
* the mutable overload, for a const array (the read path behind `x[i, j]` on a const value).403
* @param ixs one (possibly negative) coordinate per dimension.404
* @return a const reference into the backing buffer.405
* @complexity O(ndim).406
* @alloc none.407
* @test CheatahNDArray.SubscriptReadWrite408
* @crtest LangFeatures.NdarraySubscript409
*/410
template <typename... Ix>411
requires((Subscript<Ix> && ...) && sizeof...(Ix) > 0)412
const T& item_ref(Ix... ixs) const {413
return (*data_)[item_pos(ixs...)];414
}416
/// 1-D subscript: `mask[i] = v` assignment and raw element writes. Accepts an integer or a scoped417
/// enum column label (see @ref Subscript), so `q[QuatCol::W] = v` writes column W of a 1-D row.418
/// @param i the element index (negative counts from the end, Python-style).419
/// @return a mutable reference to element @p i.420
/// @complexity O(1). @alloc none.421
/// @test CheatahNDArray.SubscriptReadWrite422
template <Subscript Ix>423
T& operator[](Ix i) {424
return item_ref(subscript_index(i));425
}427
/// The buffer position of a (possibly negative) multi-index — rank and bounds are checked here,428
/// once, for both `item_ref` overloads.429
/// @param ixs one index per dimension (negative counts from the end, Python-style).430
/// @return the flat position in the buffer of that element.431
/// @complexity O(ndim). @alloc none.432
/// @test CheatahNDArray.SubscriptReadWrite433
template <typename... Ix>434
requires((Subscript<Ix> && ...) && sizeof...(Ix) > 0)435
[[nodiscard]] std::size_t item_pos(Ix... ixs) const {436
const std::array<long long, sizeof...(Ix)> raw{subscript_index(ixs)...};437
if (raw.size() != shape_.size())438
throw std::out_of_range("ndarray subscript: wrong number of indices");439
std::size_t pos = offset_;440
for (std::size_t d = 0; d < raw.size(); ++d) {441
long long i = raw[d];442
const auto n = static_cast<long long>(shape_[d]);443
if (i < 0) i += n;444
if (i < 0 || i >= n) throw std::out_of_range("ndarray subscript out of range");445
pos += static_cast<std::size_t>(i) * strides_[d];446
}447
return pos;448
}450
/**451
* Writable element at @p index — the mutable companion to @ref at, used to write through a452
* view (a slice assignment addresses elements it does not own outright).453
* @param index one coordinate per dimension.454
* @return a reference to the element.455
* @complexity O(ndim).456
* @alloc none.457
* @test CheatahNDArray.SliceAssignCopiesIn458
* @systest StdlibE2E.Ndarray459
*/460
T& at_ref(const std::vector<std::size_t>& index) {461
if (index.size() != shape_.size()) {462
throw std::runtime_error("ndarray: index has the wrong number of dimensions");463
}464
auto off = static_cast<std::ptrdiff_t>(offset_);465
for (std::size_t i = 0; i < index.size(); ++i) {466
if (index[i] >= shape_[i]) throw std::runtime_error("ndarray: index out of range");467
off += static_cast<std::ptrdiff_t>(index[i]) * strides_[i];468
}469
return (*data_)[static_cast<std::size_t>(off)];470
}471
/**472
* Read one element by a full multidimensional index (one component per axis), resolved via473
* the array's strides. Computes the flat buffer position as474
* `offset + sum(index[i] * strides[i])`, so it correctly resolves views (including475
* broadcast dims with stride 0). Bounds- and rank-checked.476
* @param index the per-axis indices; its size must equal the array's rank.477
* @return a copy of the addressed element.478
* @throws std::runtime_error if @p index has the wrong rank or any component is out of range.479
* @complexity O(rank).480
* @alloc none.481
* @test CheatahNDArray.AtVectorRankAndRangeErrors,482
* CheatahNDArray.RejectsMaliciousShapesAndIndices483
*/484
T at(const std::vector<std::size_t>& index) const { // element via strides485
// Bounds-check: a wrong-rank or out-of-range index would otherwise compute an486
// offset outside the backing buffer (out-of-bounds read).487
if (index.size() != shape_.size()) {488
throw std::runtime_error("ndarray: index has the wrong number of dimensions");489
}490
auto off = static_cast<std::ptrdiff_t>(offset_);491
for (std::size_t i = 0; i < index.size(); ++i) {492
if (index[i] >= shape_[i]) throw std::runtime_error("ndarray: index out of range");493
off += static_cast<std::ptrdiff_t>(index[i]) * strides_[i];494
}495
return (*data_)[static_cast<std::size_t>(off)];496
}497
/**498
* The shared backing buffer.499
* @return reference to the element buffer shared_ptr.500
* @complexity O(1).501
* @alloc none.502
* @test CheatahNDArray.BroadcastingAdd503
*/504
const std::shared_ptr<buffer_t<T>>& buffer() const { return data_; }505
/**506
* The flat offset into the buffer where this view starts.507
* @return the offset.508
* @complexity O(1).509
* @alloc none.510
* @test CheatahNDArray.BroadcastTo511
*/512
std::size_t offset() const { return offset_; }514
/**515
* Python-style text rendering, e.g. `"[[1, 2], [3, 4]]"` — the `str()` member, the516
* same full form `to_string` and `operator<<` produce (`io.print` reaches an array517
* through @ref cheatah_pretty_print and `operator<<`, not this hook). Defers to518
* the free `to_string`; defined out-of-line below, where `to_string` is declared.519
* @return the array formatted as nested brackets.520
* @complexity O(n) in the element count.521
* @alloc allocates the result string.522
* @test CheatahNDArray.PrettyPrintAbbreviatesLarge523
*/524
std::string str() const;526
/**527
* Pretty-print hook used by `io.print`: renders the array in nested-bracket form but528
* ABBREVIATES a large array with `...` (numpy-style edge items), so printing a big array529
* stays readable. `io.rprint` (and `str()`/`operator<<`) keep the FULL untruncated form —530
* slice the array and `rprint` the subset to see everything.531
* @param os destination stream.532
* @param indent unused (arrays are self-delimiting with brackets); present so `io.print`533
* detects this hook uniformly with struct pretty-printers.534
* @complexity O(size) — an abbreviated axis is still stepped over; only its edge items are formatted.535
* @alloc allocates intermediate strings.536
* @test CheatahNDArray.PrettyPrintAbbreviatesLarge537
* @crtest NdarrayCompileRun.PrintAbbreviatesLargeArray538
*/539
void cheatah_pretty_print(std::ostream& os, long long indent) const;541
private:542
std::shared_ptr<buffer_t<T>> data_;543
std::vector<std::size_t> shape_;544
std::vector<std::ptrdiff_t> strides_; // element strides545
std::size_t offset_ = 0;546
};548
/// The default ndarray element type is `double` — `NDArray` names that549
/// instantiation (the std::string ↔ std::basic_string<char> pattern), so existing550
/// code and the linalg routines keep working unchanged.551
using NDArray = basic_ndarray<double>;553
/**554
* The broadcast result shape of two shapes (NumPy rules).555
*556
* Aligns the shapes from the trailing (rightmost) dimension, treating missing557
* leading dims as 1; each output dim is the non-1 input dim, and two unequal dims558
* that are both not 1 are incompatible and throw.559
* @param a first shape.560
* @param b second shape.561
* @return the broadcast shape (trailing-aligned).562
* @complexity O(max(ndim)).563
* @alloc allocates the small result vector; throws if the shapes are incompatible.564
* @test CheatahNDArray.BroadcastShapeRules565
* @systest StdlibE2E.Ndarray566
*/567
std::vector<std::size_t> broadcast_shapes(const std::vector<std::size_t>& a,568
const std::vector<std::size_t>& b);570
// ==========================================================================571
// Templated free functions. The element type T is deduced from the data572
// (array([1,2,3]) -> long long, array([1.0,…]) -> double); every op is573
// constrained by Numeric. Elementwise ops vectorize via the std::execution574
// policies (declarative SIMD); broadcasting/strided views fall back to a575
// C-order scalar walk. NDArray (= basic_ndarray<double>) is the default.576
// ==========================================================================578
/**579
* Whether @p a is a contiguous C-order block (no broadcast/stride-0/permuted view),580
* so its elements live consecutively from `offset()` and can be walked flatly.581
* @param a the array (or view) to test.582
* @return true if @p a's strides are the C-order strides for its shape.583
* @complexity O(ndim).584
* @alloc none.585
* @test CheatahNDArray.BroadcastingAdd586
* @systest StdlibE2E.Ndarray587
*/588
template <Element T>589
inline bool is_contiguous(const basic_ndarray<T>& a) {590
// C-order contiguity WITHOUT materializing the reference strides: walk the dims591
// back-to-front and check each stride equals the running size product. The old592
// `strides() == contiguous_strides(shape())` heap-allocated a vector on every call593
// — a fixed cost that dominated small-n reductions (e.g. dot, where it ran twice594
// per call). Same result as the comparison, zero allocation. O(ndim).595
const std::vector<std::size_t>& shape = a.shape();596
const std::vector<std::ptrdiff_t>& strides = a.strides();597
std::ptrdiff_t expect = 1;598
for (std::size_t i = shape.size(); i-- > 0;) {599
if (strides[i] != expect) return false;600
expect *= static_cast<std::ptrdiff_t>(shape[i]);601
}602
return true;603
}605
/**606
* A zero-copy view of @p a stretched to @p target (size-1 / missing dims get stride 0).607
* @param a source array.608
* @param target the shape to stretch to.609
* @return a VIEW sharing @p a's buffer (no element copy).610
* @complexity O(rank of @p target).611
* @alloc allocates the view's shape and stride vectors, no element copy; throws if the612
* shapes are not broadcast-compatible.613
* @test CheatahNDArray.BroadcastTo614
* @systest StdlibE2E.Ndarray615
*/616
template <Element T>617
basic_ndarray<T> broadcast_to(const basic_ndarray<T>& a, const std::vector<std::size_t>& target) {618
const std::size_t n = target.size();619
if (a.ndim() > n) throw std::runtime_error("ndarray: cannot broadcast to fewer dimensions");620
std::vector<std::ptrdiff_t> ns(n, 0); // stretched / missing dims -> stride 0621
const std::size_t pad = n - a.ndim();622
for (std::size_t i = 0; i < a.ndim(); ++i) {623
const std::size_t adim = a.shape()[i];624
if (adim == target[pad + i]) {625
ns[pad + i] = a.strides()[i];626
} else if (adim != 1) {627
throw std::runtime_error("ndarray: shape not broadcastable to target");628
} // adim == 1 -> stride stays 0 (stretch)629
}630
return basic_ndarray<T>(a.buffer(), target, ns, a.offset());631
}633
// ---- factories (shapes arrive from cheatah as list[int]) ----634
/**635
* 1-D array from a list of values; the element type is the list's element type636
* (`array([1,2,3])` is integer, `array([1.0,…])` is double).637
* @param values the elements, copied into a fresh contiguous buffer.638
* @return a contiguous 1-D `basic_ndarray<T>`.639
* @complexity O(n).640
* @alloc allocates a new buffer of `values.size()` elements.641
* @test CheatahNDArray.ArrayMoveIn, CheatahNDArray.Arange642
* @crtest NdarrayCompileRun.Arange643
* @systest StdlibE2E.Ndarray644
*/645
template <Copyable T>646
requires (!detail::is_std_vector_v<T>) // a vector-of-vectors is a NESTED list (overload below)647
basic_ndarray<T> array(const std::vector<T>& values) {648
basic_ndarray<T> a = basic_ndarray<T>::uninitialized({values.size()});649
std::copy(values.begin(), values.end(), a.buffer()->begin());650
return a;651
}652
/**653
* 1-D array that MOVES its elements out of @p values into a fresh buffer (no element copy) — the654
* no-copy build path, and the ONLY `array` overload a move-only element type has. A temporary655
* `array(std::vector<T>{…})` binds here automatically; a named lvalue you want to keep uses the656
* copying overload above.657
* @param values the elements, moved into a fresh contiguous buffer (left moved-from).658
* @return a contiguous 1-D `basic_ndarray<T>`.659
* @complexity O(n).660
* @alloc allocates a new buffer of `values.size()` elements.661
* @test CheatahNDArray.ArrayMoveIn662
*/663
template <Element T>664
requires (!detail::is_std_vector_v<T>) // a vector-of-vectors is a NESTED list (overload below)665
basic_ndarray<T> array(std::vector<T>&& values) {666
basic_ndarray<T> a = basic_ndarray<T>::uninitialized({values.size()});667
std::vector<T> src = std::move(values); // consume the caller's vector; its elements move into the buffer668
std::move(src.begin(), src.end(), a.buffer()->begin());669
return a;670
}671
/**672
* N-dimensional array from a **nested** list — `array([[1, 2], [3, 4]])` is 2-D,673
* `array([[[1],[2]],[[3],[4]]])` is 3-D, and so on to any depth. The shape is read off674
* the nesting (outer list = axis 0, …) and the leaf scalar type is deduced; the list675
* must be **rectangular** (every sibling row the same length) or it throws, exactly as676
* numpy rejects a ragged array. Selected only when the argument is itself a list of677
* lists, so it never competes with the 1-D @ref array overload above.678
* @tparam V the element type of the outer list — itself a `std::vector<…>` whose leaf679
* is a @ref Field.680
* @param values the nested list (rows, planes, …), copied into a fresh C-order buffer.681
* @return a contiguous `basic_ndarray<T>` of the inferred shape.682
* @complexity O(size).683
* @alloc allocates a temporary flat vector (grown per leaf) and the result buffer; throws on a ragged list.684
* @test CheatahNDArray.NestedArrayConstruction685
* @crtest NdarrayCompileRun.NestedArray686
*/687
template <typename V>688
requires detail::is_std_vector_v<V> && Element<detail::nested_scalar_t<V>>689
basic_ndarray<detail::nested_scalar_t<V>> array(const std::vector<V>& values) {690
using T = detail::nested_scalar_t<V>;691
std::vector<std::size_t> shape;692
std::vector<T> flat;693
detail::nested_collect(values, flat, shape, 0);694
basic_ndarray<T> a = basic_ndarray<T>::uninitialized(shape);695
std::move(flat.begin(), flat.end(), a.buffer()->begin());696
return a;697
}698
/**699
* `array({1, 2, 3})` — braced-list overload (deduces T from the initializer_list,700
* which the `std::vector<T>` overload can't do directly).701
* @param values the elements as a braced list.702
* @return a contiguous 1-D `basic_ndarray<T>`.703
* @complexity O(n).704
* @alloc allocates a temporary vector and the result buffer.705
* @test CheatahNDArray.ShapeFactoriesAndReductions706
*/707
template <Copyable T>708
requires (!detail::is_std_vector_v<T>) // nested braces route to the nested-list overload709
basic_ndarray<T> array(std::initializer_list<T> values) {710
return array(std::vector<T>(values));711
}712
/**713
* 0-D scalar array (broadcasts to anything); element type deduced from @p value.714
* @param value the single element.715
* @return a 0-d `basic_ndarray<T>`.716
* @complexity O(1).717
* @alloc allocates a one-element buffer.718
* @test CheatahNDArray.ElementwiseAndScalarBroadcast719
* @crtest NdarrayCompileRun.Scalar720
* @systest StdlibE2E.Ndarray721
*/722
template <Copyable T>723
basic_ndarray<T> scalar(T value) {724
basic_ndarray<T> a; // 0-d725
a.buffer()->assign(1, value);726
return a;727
}728
/**729
* Array of @p shape filled with 0 (a `double` array by default; rejects negatives).730
* @param shape the dimensions (signed; throws on a negative).731
* @return a zero-filled `NDArray`.732
* @complexity O(size).733
* @alloc allocates a new buffer; throws on negative/overflowing dims.734
* @test CheatahNDArray.ShapeFactoriesAndReductions735
* @crtest NdarrayCompileRun.Zeros736
* @systest StdlibE2E.Ndarray737
*/738
inline NDArray zeros(const std::vector<long long>& shape) {739
return NDArray(detail::to_size(shape), 0.0);740
}741
/**742
* Array of @p shape filled with 1 (a `double` array by default; rejects negatives).743
* @param shape the dimensions (signed; throws on a negative).744
* @return a one-filled `NDArray`.745
* @complexity O(size).746
* @alloc allocates a new buffer; throws on negative/overflowing dims.747
* @test CheatahNDArray.ShapeFactoriesAndReductions748
* @crtest NdarrayCompileRun.Ones749
* @systest StdlibE2E.Ndarray750
*/751
inline NDArray ones(const std::vector<long long>& shape) {752
return NDArray(detail::to_size(shape), 1.0);753
}754
/**755
* Array of @p shape filled with @p value; element type deduced from @p value.756
* @param shape the dimensions (signed; throws on a negative).757
* @param value the fill value (its type is the array's element type).758
* @return a filled `basic_ndarray<T>`.759
* @complexity O(size).760
* @alloc allocates a new buffer; throws on negative/overflowing dims.761
* @test CheatahNDArray.RejectsMaliciousShapesAndIndices762
* @crtest NdarrayCompileRun.Full763
* @systest StdlibE2E.Ndarray764
*/765
template <Copyable T>766
basic_ndarray<T> full(const std::vector<long long>& shape, T value) {767
return basic_ndarray<T>(detail::to_size(shape), value);768
}769
/**770
* A fresh array with the SAME shape and element type as @p a, filled with @p value771
* (≈ `numpy.full_like`). The companion `zeros_like` / `ones_like` default the fill.772
* @param a the array whose shape and element type to mirror.773
* @param value the fill value.774
* @return a same-shape `basic_ndarray<T>` filled with @p value.775
* @complexity O(size).776
* @alloc allocates a new buffer.777
* @test CheatahNDArray.LikeFactories778
*/779
template <Copyable T>780
basic_ndarray<T> full_like(const basic_ndarray<T>& a, T value) {781
return basic_ndarray<T>(a.shape(), value);782
}783
/**784
* A zero-filled array with the SAME shape and element type as @p a (≈ `numpy.zeros_like`) —785
* the idiomatic way to allocate a matching gradient/velocity/scratch buffer for an existing array.786
* @param a the array whose shape and element type to mirror.787
* @return a same-shape `basic_ndarray<T>` of zeros.788
* @complexity O(size).789
* @alloc allocates a new buffer.790
* @test CheatahNDArray.LikeFactories791
*/792
template <Copyable T>793
basic_ndarray<T> zeros_like(const basic_ndarray<T>& a) {794
return basic_ndarray<T>(a.shape(), T{});795
}796
/**797
* A one-filled array with the SAME shape and element type as @p a (≈ `numpy.ones_like`).798
* @param a the array whose shape and element type to mirror.799
* @return a same-shape `basic_ndarray<T>` of ones.800
* @complexity O(size).801
* @alloc allocates a new buffer.802
* @test CheatahNDArray.LikeFactories803
*/804
template <Copyable T>805
basic_ndarray<T> ones_like(const basic_ndarray<T>& a) {806
return basic_ndarray<T>(a.shape(), T{1});807
}808
/**809
* 1-D range `[start, stop)` stepping by @p step; element type deduced from the args.810
* @param start first value.811
* @param stop exclusive bound.812
* @param step increment (throws if zero); a step pointing away from @p stop yields empty.813
* @return a 1-D `basic_ndarray<T>` of the generated values.814
* @complexity O(count).815
* @alloc allocates a new buffer (built via a growing temporary vector, then copied);816
* throws if @p step is zero.817
* @test CheatahNDArray.Arange818
* @crtest NdarrayCompileRun.Arange819
* @systest StdlibE2E.Ndarray820
*/821
template <Numeric T>822
basic_ndarray<T> arange(T start, T stop, T step) {823
if (step == T{}) throw std::runtime_error("ndarray: arange step must be non-zero");824
std::vector<T> v;825
if constexpr (std::floating_point<T>) {826
// Integer induction with x = start + i*step — numpy computes arange the same way.827
// A floating-point loop counter (cert-flp30-c) accumulates rounding error every828
// pass; the multiply form keeps each element one rounding away from exact. The829
// bound check per element preserves the exclusive-stop semantics at the boundary.830
const double span = (static_cast<double>(stop) - static_cast<double>(start)) /831
static_cast<double>(step);832
const std::size_t count = span > 0.0 ? static_cast<std::size_t>(std::ceil(span)) : 0;833
v.reserve(count);834
for (std::size_t i = 0; i < count; ++i) {835
const T x = start + static_cast<T>(i) * step;836
if ((step > T{}) ? (x < stop) : (x > stop)) v.push_back(x);837
}838
} else {839
for (T x = start; (step > T{}) ? (x < stop) : (x > stop); x += step) v.push_back(x);840
}841
return array(v);842
}843
/**844
* Reshape @p a to @p shape (same element count); reads in C-order so views/broadcasts845
* are flattened into a fresh contiguous buffer (a copy, not an alias).846
* @param a source array.847
* @param shape the new dimensions (signed; throws on a negative).848
* @return a new contiguous `basic_ndarray<T>` with the data in C-order.849
* @complexity O(size).850
* @alloc allocates a new buffer; throws on size mismatch or negative dims.851
* @test CheatahNDArray.BroadcastingAdd, CheatahNDArray.ReshapeSizeMismatchThrows852
* @crtest NdarrayCompileRun.Reshape853
* @systest StdlibE2E.Ndarray854
*/855
template <Copyable T>856
basic_ndarray<T> reshape(const basic_ndarray<T>& a, const std::vector<long long>& shape) {857
const std::vector<std::size_t> ns = detail::to_size(shape);858
if (detail::product(ns) != a.size()) {859
throw std::runtime_error("ndarray: cannot reshape, size mismatch");860
}861
basic_ndarray<T> out = basic_ndarray<T>::uninitialized(ns); // every element is written below862
auto& buf = *out.buffer();863
// Contiguous source (the common case — e.g. reshaping a freshly built array): copy864
// the flat block in one shot instead of walking a per-element bounds-checked865
// odometer.866
if (is_contiguous(a)) {867
const T* src = a.buffer()->data() + a.offset();868
std::copy(src, src + a.size(), buf.begin());869
return out;870
}871
std::vector<std::size_t> idx(a.ndim(), 0);872
std::size_t flat = 0;873
for (;;) { // a 0-d array still yields its one element874
buf[flat++] = a.at(idx);875
if (a.ndim() == 0 || !detail::next_index(idx, a.shape())) break;876
}877
return out;878
}879
/**880
* Convert @p a to a new array with element type @p U — numpy's `a.astype(dtype)`. Every element881
* is `static_cast` into @p U, so this is the way to build a NARROW-element array (a smaller memory882
* footprint): `array([1,2,3]).astype(i16)` is a `basic_ndarray<std::int16_t>` — 2 bytes/element,883
* not 8. Reads @p a in C-order (a view/broadcast is flattened into a fresh contiguous buffer — a884
* copy, never an alias), same shape out as in. Widening is exact; narrowing truncates/wraps at the885
* target width (as in C / a numpy fixed dtype). Constrained to conversions that actually exist886
* (`convertible_to`), so e.g. complex→real fails with a clear concept error, not template spam.887
* @tparam U the destination element type (the only type spelled at the call site).888
* @param a source array (any @ref Field element type convertible to @p U).889
* @return a fresh contiguous `basic_ndarray<U>` of @p a's shape.890
* @complexity O(size).891
* @alloc allocates the result buffer.892
* @test CheatahNDArray.AstypeNarrowsAndWidens893
* @crtest NdarrayCompileRun.Astype894
* @systest StdlibE2E.Ndarray895
*/896
template <Field U, Field T>897
requires std::convertible_to<T, U>898
basic_ndarray<U> astype(const basic_ndarray<T>& a) {899
basic_ndarray<U> out = basic_ndarray<U>::uninitialized(a.shape()); // every element is written below900
auto& buf = *out.buffer();901
if (is_contiguous(a)) { // contiguous source: one straight cast pass, no odometer902
const T* src = a.buffer()->data() + a.offset();903
// NOLINT below: an i8→wider astype must sign-extend (numpy dtype semantics) — these904
// are small integers, not characters, so the signed-char-misuse hazard does not apply.905
for (std::size_t i = 0; i < a.size(); ++i) buf[i] = static_cast<U>(src[i]); // NOLINT(bugprone-signed-char-misuse,cert-str34-c)906
return out;907
}908
std::vector<std::size_t> idx(a.ndim(), 0);909
std::size_t flat = 0;910
for (;;) { // a 0-d array still yields its one element911
buf[flat++] = static_cast<U>(a.at(idx)); // NOLINT(bugprone-signed-char-misuse,cert-str34-c): same sign-extension intent as above912
if (a.ndim() == 0 || !detail::next_index(idx, a.shape())) break;913
}914
return out;915
}917
// ---- element-wise ops (broadcasting, vectorized) ----918
/// @cond INTERNAL919
/// implementation plumbing / compiler-selected reuse overloads (README documents the public forms)920
/**921
* Broadcast @p a and @p b to their common shape and apply @p op elementwise into a922
* fresh contiguous result. Fast path: when both operands are contiguous, a flat923
* `std::transform` under the `unseq` policy (SIMD); otherwise a C-order scalar walk.924
* @param a first operand.925
* @param b second operand.926
* @param op the binary operation applied to corresponding elements.927
* @return `op(a, b)` broadcast to the common shape; throws if shapes don't broadcast.928
* @complexity O(size of result).929
* @alloc allocates the result buffer.930
* @test CheatahNDArray.BroadcastingAdd931
* @systest StdlibE2E.Ndarray932
*/933
template <Field T, typename Op>934
basic_ndarray<T> binary_op(const basic_ndarray<T>& a, const basic_ndarray<T>& b, Op op) {935
const std::vector<std::size_t> rshape = broadcast_shapes(a.shape(), b.shape());936
basic_ndarray<T> out = basic_ndarray<T>::uninitialized(rshape); // every element written below937
auto& obuf = *out.buffer();938
// Scalar fast paths: `array ⊕ scalar` (or the reverse) is by far the most common939
// broadcast, and the general strided walk below does a bounds-checked at() per940
// element (no SIMD). When the other operand is a single value over a contiguous941
// full-shape array, it's a flat loop we hand to the unseq transform so it vectorizes942
// the same way the array⊕array path does.943
if (b.size() == 1 && a.shape() == rshape && is_contiguous(a)) {944
const T s = (*b.buffer())[b.offset()];945
const auto first = a.buffer()->begin() + a.offset();946
std::transform(CHEATAH_UNSEQ first, first + obuf.size(), obuf.begin(),947
[s, op](T x) { return op(x, s); });948
return out;949
}950
if (a.size() == 1 && b.shape() == rshape && is_contiguous(b)) {951
const T s = (*a.buffer())[a.offset()];952
const auto first = b.buffer()->begin() + b.offset();953
std::transform(CHEATAH_UNSEQ first, first + obuf.size(), obuf.begin(),954
[s, op](T x) { return op(s, x); });955
return out;956
}957
const basic_ndarray<T> av = broadcast_to(a, rshape);958
const basic_ndarray<T> bv = broadcast_to(b, rshape);959
if (is_contiguous(av) && is_contiguous(bv)) {960
const auto& abuf = *av.buffer();961
const auto& bbuf = *bv.buffer();962
std::transform(CHEATAH_UNSEQ abuf.begin() + av.offset(),963
abuf.begin() + av.offset() + obuf.size(), bbuf.begin() + bv.offset(),964
obuf.begin(), op);965
return out;966
}967
std::vector<std::size_t> idx(rshape.size(), 0);968
std::size_t flat = 0;969
for (;;) { // a 0-d result still has its one element970
obuf[flat++] = op(av.at(idx), bv.at(idx));971
if (rshape.empty() || !detail::next_index(idx, rshape)) break;972
}973
return out;974
}976
/**977
* Elementwise `out = op(a, b)` (broadcasting) into the CALLER'S buffer @p out — the user-provided-output978
* form of binary_op, NO allocation: a hot loop hands the same scratch array every call. @p out must979
* already hold the broadcast result shape and be contiguous; it MAY alias a full-shape operand (the write980
* is index-local, so `add(x, x, y)` is fine).981
* @param out the destination (mutated; must be contiguous and match the broadcast shape).982
* @param a first operand.983
* @param b second operand.984
* @param op the binary combiner.985
* @complexity O(size of out).986
* @alloc none.987
* @test CheatahNDArray.BinaryOpIntoReusesBuffer988
*/989
template <Field T, typename Op>990
void binary_op_into(basic_ndarray<T>& out, const basic_ndarray<T>& a, const basic_ndarray<T>& b, Op op) {991
const std::vector<std::size_t> rshape = broadcast_shapes(a.shape(), b.shape());992
if (out.shape() != rshape || !is_contiguous(out)) {993
throw std::invalid_argument(994
"ndarray binary op (out form): out must be contiguous and match the broadcast shape");995
}996
auto& obuf = *out.buffer();997
const auto odst = obuf.begin() + static_cast<std::ptrdiff_t>(out.offset());998
if (b.size() == 1 && a.shape() == rshape && is_contiguous(a)) { // array ⊕ scalar999
const T s = (*b.buffer())[b.offset()];1000
const auto af = a.buffer()->begin() + static_cast<std::ptrdiff_t>(a.offset());1001
std::transform(CHEATAH_UNSEQ af, af + static_cast<std::ptrdiff_t>(out.size()), odst,1002
[s, op](T x) { return op(x, s); });1003
return;1004
}1005
if (a.size() == 1 && b.shape() == rshape && is_contiguous(b)) { // scalar ⊕ array1006
const T s = (*a.buffer())[a.offset()];1007
const auto bf = b.buffer()->begin() + static_cast<std::ptrdiff_t>(b.offset());1008
std::transform(CHEATAH_UNSEQ bf, bf + static_cast<std::ptrdiff_t>(out.size()), odst,1009
[s, op](T x) { return op(s, x); });1010
return;1011
}1012
const basic_ndarray<T> av = broadcast_to(a, rshape);1013
const basic_ndarray<T> bv = broadcast_to(b, rshape);1014
if (is_contiguous(av) && is_contiguous(bv)) { // both full-shape contiguous: flat SIMD transform1015
const auto af = av.buffer()->begin() + static_cast<std::ptrdiff_t>(av.offset());1016
std::transform(CHEATAH_UNSEQ af, af + static_cast<std::ptrdiff_t>(out.size()),1017
bv.buffer()->begin() + static_cast<std::ptrdiff_t>(bv.offset()), odst, op);1018
return;1019
}1020
std::vector<std::size_t> idx(rshape.size(), 0); // strided operand fallback (still no alloc for out)1021
std::size_t flat = 0;1022
for (;;) { // a 0-d result still has its one element1023
odst[static_cast<std::ptrdiff_t>(flat++)] = op(av.at(idx), bv.at(idx));1024
if (rshape.empty() || !detail::next_index(idx, rshape)) break;1025
}1026
}1027
/// @endcond1029
// Shared elementwise combiners: ONE functor type per op, used by BOTH the allocating1030
// forms (add/sub/mul/divide) and the in-place compound operators (+=/-=/*=//=). Using a1031
// single type means `binary_op` is instantiated once per op rather than once per call1032
// site, so the in-place fallback reuses the same (already-tested) instantiation instead1033
// of a duplicate whose scalar/contiguous fast paths are unreachable through it.1034
namespace detail {1035
struct add_op { template <typename T> T operator()(T x, T y) const { return x + y; } };1036
struct sub_op { template <typename T> T operator()(T x, T y) const { return x - y; } };1037
struct mul_op { template <typename T> T operator()(T x, T y) const { return x * y; } };1038
struct div_op { template <typename T> T operator()(T x, T y) const { return x / y; } };1039
// Reversed combiners: reuse the RIGHT operand in place for non-commutative ops, i.e. compute1040
// `dst = src OP dst` so that `a - std::move(b)` / `a / std::move(b)` can write through b's buffer.1041
struct rsub_op { template <typename T> T operator()(T x, T y) const { return y - x; } };1042
struct rdiv_op { template <typename T> T operator()(T x, T y) const { return y / x; } };1043
} // namespace detail1045
/**1046
* Element-wise `a + b` with broadcasting.1047
* @param a first operand.1048
* @param b second operand.1049
* @return `a + b` broadcast to the common shape.1050
* @complexity O(size of result). @alloc allocates the result.1051
* @test CheatahNDArray.BroadcastingAdd1052
* @crtest NdarrayCompileRun.Add1053
* @systest StdlibE2E.Ndarray1054
*/1055
template <Field T>1056
basic_ndarray<T> add(const basic_ndarray<T>& a, const basic_ndarray<T>& b) {1057
return binary_op(a, b, detail::add_op{});1058
}1059
/**1060
* Element-wise `a - b` with broadcasting.1061
* @param a first operand.1062
* @param b second operand.1063
* @return `a - b` broadcast to the common shape.1064
* @complexity O(size of result). @alloc allocates the result.1065
* @test CheatahNDArray.ElementwiseAndScalarBroadcast1066
* @crtest NdarrayCompileRun.Sub1067
* @systest StdlibE2E.Ndarray1068
*/1069
template <Field T>1070
basic_ndarray<T> sub(const basic_ndarray<T>& a, const basic_ndarray<T>& b) {1071
return binary_op(a, b, detail::sub_op{});1072
}1073
/**1074
* Element-wise `a * b` with broadcasting.1075
* @param a first operand.1076
* @param b second operand.1077
* @return `a * b` broadcast to the common shape.1078
* @complexity O(size of result). @alloc allocates the result.1079
* @test CheatahNDArray.ElementwiseAndScalarBroadcast1080
* @crtest NdarrayCompileRun.Mul1081
* @systest StdlibE2E.Ndarray1082
*/1083
template <Field T>1084
basic_ndarray<T> mul(const basic_ndarray<T>& a, const basic_ndarray<T>& b) {1085
return binary_op(a, b, detail::mul_op{});1086
}1087
/**1088
* Element-wise `a / b` with broadcasting (an integer element type does integer division).1089
* @param a numerator.1090
* @param b denominator (float division follows IEEE-754: /0 yields inf/nan, no throw).1091
* @return `a / b` broadcast to the common shape.1092
* @complexity O(size of result). @alloc allocates the result.1093
* @test CheatahNDArray.ElementwiseAndScalarBroadcast1094
* @crtest NdarrayCompileRun.Divide1095
* @systest StdlibE2E.Ndarray1096
*/1097
template <Field T>1098
basic_ndarray<T> divide(const basic_ndarray<T>& a, const basic_ndarray<T>& b) {1099
return binary_op(a, b, detail::div_op{});1100
}1102
/// @cond INTERNAL1103
/// implementation plumbing / compiler-selected reuse overloads (README documents the public forms)1104
// ---- user-provided-output forms: write into the caller's buffer, NO allocation (see binary_op_into) ----1105
/**1106
* Element-wise `a + b` into the caller's buffer @p out (out FIRST) — the buffer-reuse overload, so a1107
* hot loop hands the same scratch every call. @p out must be contiguous with the broadcast shape; it1108
* may alias a full-shape operand (the write is index-local).1109
* @param out destination, overwritten.1110
* @param a,b operands (broadcastable to @p out's shape).1111
* @complexity O(size of result). @alloc none.1112
* @test CheatahNDArray.BinaryOpIntoReusesBuffer1113
*/1114
template <Field T>1115
void add(basic_ndarray<T>& out, const basic_ndarray<T>& a, const basic_ndarray<T>& b) {1116
binary_op_into(out, a, b, detail::add_op{});1117
}1118
/**1119
* Element-wise `a - b` into @p out (out FIRST), no allocation. @see add(out, a, b).1120
* @param out destination, overwritten.1121
* @param a,b operands (broadcastable to @p out's shape).1122
*/1123
template <Field T>1124
void sub(basic_ndarray<T>& out, const basic_ndarray<T>& a, const basic_ndarray<T>& b) {1125
binary_op_into(out, a, b, detail::sub_op{});1126
}1127
/**1128
* Element-wise `a * b` into @p out (out FIRST), no allocation. @see add(out, a, b).1129
* @param out destination, overwritten.1130
* @param a,b operands (broadcastable to @p out's shape).1131
*/1132
template <Field T>1133
void mul(basic_ndarray<T>& out, const basic_ndarray<T>& a, const basic_ndarray<T>& b) {1134
binary_op_into(out, a, b, detail::mul_op{});1135
}1136
/**1137
* Element-wise `a / b` into @p out (out FIRST), no allocation. @see add(out, a, b).1138
* @param out destination, overwritten.1139
* @param a,b operands (broadcastable to @p out's shape).1140
*/1141
template <Field T>1142
void divide(basic_ndarray<T>& out, const basic_ndarray<T>& a, const basic_ndarray<T>& b) {1143
binary_op_into(out, a, b, detail::div_op{});1144
}1145
/// @endcond1147
// ---- Infix operators & in-place compound assignment ------------------------1148
// cheatah lowers `a + b` / `a * 2.0` / `a += b` on ndarrays straight to these1149
// C++ operators. Infix forms are the elementwise free functions (broadcasting1150
// included); a bare arithmetic scalar on either side is wrapped via scalar().1151
// The compound forms mutate the LEFT OPERAND'S BUFFER IN PLACE on the common1152
// (contiguous) layout — no allocation, so a hot loop can reuse one array for1153
// an entire run — falling back to the allocating elementwise path only for1154
// non-contiguous views or true broadcasts.1156
/// @cond INTERNAL1157
/// implementation plumbing / compiler-selected reuse overloads (README documents the public forms)1158
/**1159
* In-place elementwise update `a = op(a, b)`, writing through @p a's buffer.1160
* Contiguous @p a with a same-shape contiguous or single-element @p b runs as1161
* a flat vectorizable transform with NO allocation; anything else falls back1162
* to the allocating binary_op and rebinds @p a to the result.1163
* @param a the destination array (mutated).1164
* @param b the right operand (same shape, or a single element).1165
* @param op the elementwise combiner.1166
* @complexity O(size of @p a).1167
* @alloc none on the contiguous fast path; one result array on the fallback.1168
* @test CheatahNDArray.CompoundAssignInPlace1169
* @crtest LangFeatures.NdarrayOperators1170
* @systest StdlibE2E.Ndarray1171
*/1172
template <typename T, typename Op>1173
void compound_apply(basic_ndarray<T>& a, const basic_ndarray<T>& b, Op op) {1174
if (is_contiguous(a) && (b.size() == 1 || b.shape() == a.shape())) {1175
auto& abuf = *a.buffer();1176
const std::size_t n = a.size();1177
const auto first = abuf.begin() + static_cast<std::ptrdiff_t>(a.offset());1178
if (b.size() == 1) {1179
const T s = (*b.buffer())[b.offset()];1180
std::transform(CHEATAH_UNSEQ first, first + static_cast<std::ptrdiff_t>(n), first,1181
[s, op](T x) { return op(x, s); });1182
return;1183
}1184
if (is_contiguous(b)) {1185
const auto& bbuf = *b.buffer();1186
std::transform(CHEATAH_UNSEQ first, first + static_cast<std::ptrdiff_t>(n),1187
bbuf.begin() + static_cast<std::ptrdiff_t>(b.offset()), first, op);1188
return;1189
}1190
}1191
a = binary_op(a, b, op); // broadcast / non-contiguous fallback1192
}1193
/// @endcond1195
/// @cond INTERNAL1196
/// implementation plumbing / compiler-selected reuse overloads (README documents the public forms)1197
// ---- rvalue-reuse ("move") forms -----------------------------------------------------------------1198
// "Copy vs move" for ndarray math. `a + b` on NAMED (lvalue) arrays MUST allocate — it can't clobber1199
// `a`, which the caller still holds. But when the LEFT operand is an RVALUE — a temporary the caller1200
// has already given up: the `a + b` inside a chain `a + b + c`, or an explicit `std::move(a)` — these1201
// compute IN PLACE into that buffer and move it out: NO allocation. Selected by value category, so a1202
// buffer is only ever reused when it is safe to (no flag, no surprise mutation). Reuses the in-place1203
// compound_apply, so the same contiguous fast path / broadcast fallback / tests apply.1204
/**1205
* Element-wise `a + b` reusing the expiring left operand @p a in place (no allocation).1206
* @param a the expiring left operand; computed into and moved out.1207
* @param b the right operand (broadcastable to @p a's shape).1208
* @return the sum, in @p a's reused buffer.1209
* @complexity O(size of @p a). @alloc none on the contiguous fast path.1210
*/1211
template <Field T>1212
basic_ndarray<T> add(basic_ndarray<T>&& a, const basic_ndarray<T>& b) {1213
compound_apply(a, b, detail::add_op{});1214
return std::move(a);1215
}1216
/**1217
* Element-wise `a - b` reusing the expiring left operand @p a in place (no allocation).1218
* @param a the expiring left operand; computed into and moved out.1219
* @param b the right operand (broadcastable to @p a's shape).1220
* @return the difference, in @p a's reused buffer.1221
* @complexity O(size of @p a). @alloc none on the contiguous fast path.1222
*/1223
template <Field T>1224
basic_ndarray<T> sub(basic_ndarray<T>&& a, const basic_ndarray<T>& b) {1225
compound_apply(a, b, detail::sub_op{});1226
return std::move(a);1227
}1228
/**1229
* Element-wise `a * b` reusing the expiring left operand @p a in place (no allocation).1230
* @param a the expiring left operand; computed into and moved out.1231
* @param b the right operand (broadcastable to @p a's shape).1232
* @return the product, in @p a's reused buffer.1233
* @complexity O(size of @p a). @alloc none on the contiguous fast path.1234
*/1235
template <Field T>1236
basic_ndarray<T> mul(basic_ndarray<T>&& a, const basic_ndarray<T>& b) {1237
compound_apply(a, b, detail::mul_op{});1238
return std::move(a);1239
}1240
/**1241
* Element-wise `a / b` reusing the expiring left operand @p a in place (no allocation).1242
* @param a the expiring left operand; computed into and moved out.1243
* @param b the right operand (broadcastable to @p a's shape).1244
* @return the quotient, in @p a's reused buffer.1245
* @complexity O(size of @p a). @alloc none on the contiguous fast path.1246
*/1247
template <Field T>1248
basic_ndarray<T> divide(basic_ndarray<T>&& a, const basic_ndarray<T>& b) {1249
compound_apply(a, b, detail::div_op{});1250
return std::move(a);1251
}1253
// Right-operand reuse: when only the RIGHT operand is the expiring temporary, compute through ITS1254
// buffer instead. `+`/`*` are commutative so `op(b, a)` is the same value; `-`/`/` use the reversed1255
// combiners (`b = a OP b`). This makes `a + std::move(b)` reuse a buffer exactly like1256
// `std::move(a) + b` — the two are symmetric, neither allocates.1257
/**1258
* Element-wise `a + b` reusing the expiring right operand @p b in place (no allocation).1259
* @param a the left operand (a const lvalue the caller keeps).1260
* @param b the expiring right operand; computed into and moved out.1261
* @return the sum, in @p b's reused buffer.1262
* @complexity O(size of @p b). @alloc none on the contiguous fast path.1263
*/1264
template <Field T>1265
basic_ndarray<T> add(const basic_ndarray<T>& a, basic_ndarray<T>&& b) {1266
compound_apply(b, a, detail::add_op{});1267
return std::move(b);1268
}1269
/**1270
* Element-wise `a - b` reusing the expiring right operand @p b in place via the reversed combiner1271
* `b = a - b` (no allocation).1272
* @param a the left operand (a const lvalue the caller keeps).1273
* @param b the expiring right operand; computed into and moved out.1274
* @return the difference, in @p b's reused buffer.1275
* @complexity O(size of @p b). @alloc none on the contiguous fast path.1276
*/1277
template <Field T>1278
basic_ndarray<T> sub(const basic_ndarray<T>& a, basic_ndarray<T>&& b) {1279
compound_apply(b, a, detail::rsub_op{});1280
return std::move(b);1281
}1282
/**1283
* Element-wise `a * b` reusing the expiring right operand @p b in place (no allocation).1284
* @param a the left operand (a const lvalue the caller keeps).1285
* @param b the expiring right operand; computed into and moved out.1286
* @return the product, in @p b's reused buffer.1287
* @complexity O(size of @p b). @alloc none on the contiguous fast path.1288
*/1289
template <Field T>1290
basic_ndarray<T> mul(const basic_ndarray<T>& a, basic_ndarray<T>&& b) {1291
compound_apply(b, a, detail::mul_op{});1292
return std::move(b);1293
}1294
/**1295
* Element-wise `a / b` reusing the expiring right operand @p b in place via the reversed combiner1296
* `b = a / b` (no allocation).1297
* @param a the numerator (a const lvalue the caller keeps).1298
* @param b the expiring denominator; computed into and moved out.1299
* @return the quotient, in @p b's reused buffer.1300
* @complexity O(size of @p b). @alloc none on the contiguous fast path.1301
*/1302
template <Field T>1303
basic_ndarray<T> divide(const basic_ndarray<T>& a, basic_ndarray<T>&& b) {1304
compound_apply(b, a, detail::rdiv_op{});1305
return std::move(b);1306
}1308
// Both operands expiring: prefer reusing the LEFT (matches the chain `a + b + c`, where the left is1309
// the running accumulator). Disambiguates the otherwise-ambiguous `std::move(a) OP std::move(b)`.1310
/**1311
* Element-wise `a + b` when both operands are expiring; reuses the left buffer @p a (no allocation).1312
* @param a the expiring left operand; reused for the result.1313
* @param b the expiring right operand.1314
* @return the sum, in @p a's reused buffer.1315
* @complexity O(size of @p a). @alloc none on the contiguous fast path.1316
*/1317
template <Field T>1318
basic_ndarray<T> add(basic_ndarray<T>&& a, basic_ndarray<T>&& b) { return add(std::move(a), b); } // NOLINT(cppcoreguidelines-rvalue-reference-param-not-moved): b is deliberately NOT consumed — the left buffer is reused; this overload exists only to disambiguate (&&, &&)1319
/**1320
* Element-wise `a - b` when both operands are expiring; reuses the left buffer @p a (no allocation).1321
* @param a the expiring left operand; reused for the result.1322
* @param b the expiring right operand.1323
* @return the difference, in @p a's reused buffer.1324
* @complexity O(size of @p a). @alloc none on the contiguous fast path.1325
*/1326
template <Field T>1327
basic_ndarray<T> sub(basic_ndarray<T>&& a, basic_ndarray<T>&& b) { return sub(std::move(a), b); } // NOLINT(cppcoreguidelines-rvalue-reference-param-not-moved): b is deliberately NOT consumed — the left buffer is reused; this overload exists only to disambiguate (&&, &&)1328
/**1329
* Element-wise `a * b` when both operands are expiring; reuses the left buffer @p a (no allocation).1330
* @param a the expiring left operand; reused for the result.1331
* @param b the expiring right operand.1332
* @return the product, in @p a's reused buffer.1333
* @complexity O(size of @p a). @alloc none on the contiguous fast path.1334
*/1335
template <Field T>1336
basic_ndarray<T> mul(basic_ndarray<T>&& a, basic_ndarray<T>&& b) { return mul(std::move(a), b); } // NOLINT(cppcoreguidelines-rvalue-reference-param-not-moved): b is deliberately NOT consumed — the left buffer is reused; this overload exists only to disambiguate (&&, &&)1337
/**1338
* Element-wise `a / b` when both operands are expiring; reuses the left buffer @p a (no allocation).1339
* @param a the expiring numerator; reused for the result.1340
* @param b the expiring denominator.1341
* @return the quotient, in @p a's reused buffer.1342
* @complexity O(size of @p a). @alloc none on the contiguous fast path.1343
*/1344
template <Field T>1345
basic_ndarray<T> divide(basic_ndarray<T>&& a, basic_ndarray<T>&& b) { return divide(std::move(a), b); } // NOLINT(cppcoreguidelines-rvalue-reference-param-not-moved): b is deliberately NOT consumed — the left buffer is reused; this overload exists only to disambiguate (&&, &&)1346
/// @endcond1348
/// Elementwise infix forms: `a + b`, `a - b`, `a * b`, `a / b` (broadcasting).1349
/**1350
* Elementwise `a + b` with broadcasting (infix form of add()).1351
* @param a first operand. @param b second operand.1352
* @return the broadcast sum (a fresh array). @complexity O(size of result). @alloc the result.1353
* @test CheatahNDArray.RvalueOperandReusesBuffer1354
*/1355
template <typename T>1356
basic_ndarray<T> operator+(const basic_ndarray<T>& a, const basic_ndarray<T>& b) { return add(a, b); }1357
/**1358
* Elementwise `a - b` with broadcasting (infix form of sub()).1359
* @param a first operand. @param b second operand.1360
* @return the broadcast difference (a fresh array). @complexity O(size of result). @alloc the result.1361
*/1362
template <typename T>1363
basic_ndarray<T> operator-(const basic_ndarray<T>& a, const basic_ndarray<T>& b) { return sub(a, b); }1364
/**1365
* Elementwise `a * b` with broadcasting (infix form of mul()).1366
* @param a first operand. @param b second operand.1367
* @return the broadcast product (a fresh array). @complexity O(size of result). @alloc the result.1368
*/1369
template <typename T>1370
basic_ndarray<T> operator*(const basic_ndarray<T>& a, const basic_ndarray<T>& b) { return mul(a, b); }1371
/**1372
* Elementwise `a / b` with broadcasting (infix form of divide()).1373
* @param a numerator. @param b denominator.1374
* @return the broadcast quotient (a fresh array). @complexity O(size of result). @alloc the result.1375
* @test CheatahNDArray.DivideInfixLvalueForm1376
*/1377
template <typename T>1378
basic_ndarray<T> operator/(const basic_ndarray<T>& a, const basic_ndarray<T>& b) { return divide(a, b); }1380
/// @cond INTERNAL1381
/// implementation plumbing / compiler-selected reuse overloads (README documents the public forms)1382
/// rvalue-reuse infix forms: whichever operand is the expiring temporary is computed into IN PLACE1383
/// (no alloc). `std::move(a) + b` reuses `a`, `a + std::move(b)` reuses `b` — symmetric. A chain1384
/// `a + b + c` allocates once (for `a + b`) instead of twice. If BOTH are temporaries the left wins.1385
/** `a + b` reusing the expiring left operand @p a in place. @param a expiring left operand. @param b right operand. @return the sum, in @p a's buffer. @complexity O(size). @alloc none on the fast path. */1386
template <typename T>1387
basic_ndarray<T> operator+(basic_ndarray<T>&& a, const basic_ndarray<T>& b) { return add(std::move(a), b); }1388
/** `a - b` reusing the expiring left operand @p a in place. @param a expiring left operand. @param b right operand. @return the difference, in @p a's buffer. @complexity O(size). @alloc none on the fast path. */1389
template <typename T>1390
basic_ndarray<T> operator-(basic_ndarray<T>&& a, const basic_ndarray<T>& b) { return sub(std::move(a), b); }1391
/** `a * b` reusing the expiring left operand @p a in place. @param a expiring left operand. @param b right operand. @return the product, in @p a's buffer. @complexity O(size). @alloc none on the fast path. */1392
template <typename T>1393
basic_ndarray<T> operator*(basic_ndarray<T>&& a, const basic_ndarray<T>& b) { return mul(std::move(a), b); }1394
/** `a / b` reusing the expiring left operand @p a in place. @param a expiring numerator. @param b denominator. @return the quotient, in @p a's buffer. @complexity O(size). @alloc none on the fast path. */1395
template <typename T>1396
basic_ndarray<T> operator/(basic_ndarray<T>&& a, const basic_ndarray<T>& b) { return divide(std::move(a), b); }1397
/** `a + b` reusing the expiring right operand @p b in place. @param a left operand. @param b expiring right operand. @return the sum, in @p b's buffer. @complexity O(size). @alloc none on the fast path. */1398
template <typename T>1399
basic_ndarray<T> operator+(const basic_ndarray<T>& a, basic_ndarray<T>&& b) { return add(a, std::move(b)); }1400
/** `a - b` reusing the expiring right operand @p b in place. @param a left operand. @param b expiring right operand. @return the difference, in @p b's buffer. @complexity O(size). @alloc none on the fast path. */1401
template <typename T>1402
basic_ndarray<T> operator-(const basic_ndarray<T>& a, basic_ndarray<T>&& b) { return sub(a, std::move(b)); }1403
/** `a * b` reusing the expiring right operand @p b in place. @param a left operand. @param b expiring right operand. @return the product, in @p b's buffer. @complexity O(size). @alloc none on the fast path. */1404
template <typename T>1405
basic_ndarray<T> operator*(const basic_ndarray<T>& a, basic_ndarray<T>&& b) { return mul(a, std::move(b)); }1406
/** `a / b` reusing the expiring right operand @p b in place. @param a numerator. @param b expiring denominator. @return the quotient, in @p b's buffer. @complexity O(size). @alloc none on the fast path. */1407
template <typename T>1408
basic_ndarray<T> operator/(const basic_ndarray<T>& a, basic_ndarray<T>&& b) { return divide(a, std::move(b)); }1409
/** `a + b` with both operands expiring; reuses the left buffer @p a. @param a expiring left operand. @param b expiring right operand. @return the sum, in @p a's buffer. @complexity O(size). @alloc none on the fast path. */1410
template <typename T>1411
basic_ndarray<T> operator+(basic_ndarray<T>&& a, basic_ndarray<T>&& b) { return add(std::move(a), std::move(b)); }1412
/** `a - b` with both operands expiring; reuses the left buffer @p a. @param a expiring left operand. @param b expiring right operand. @return the difference, in @p a's buffer. @complexity O(size). @alloc none on the fast path. */1413
template <typename T>1414
basic_ndarray<T> operator-(basic_ndarray<T>&& a, basic_ndarray<T>&& b) { return sub(std::move(a), std::move(b)); }1415
/** `a * b` with both operands expiring; reuses the left buffer @p a. @param a expiring left operand. @param b expiring right operand. @return the product, in @p a's buffer. @complexity O(size). @alloc none on the fast path. */1416
template <typename T>1417
basic_ndarray<T> operator*(basic_ndarray<T>&& a, basic_ndarray<T>&& b) { return mul(std::move(a), std::move(b)); }1418
/** `a / b` with both operands expiring; reuses the left buffer @p a. @param a expiring numerator. @param b expiring denominator. @return the quotient, in @p a's buffer. @complexity O(size). @alloc none on the fast path. */1419
template <typename T>1420
basic_ndarray<T> operator/(basic_ndarray<T>&& a, basic_ndarray<T>&& b) { return divide(std::move(a), std::move(b)); }1421
/// @endcond1423
/// Scalar infix forms, either side: `a * 2.0`, `0.5 * a`, `a + 1`, `1.0 / a`, ...1424
/// The scalar converts to the array's element type (int literals work on float arrays). These go1425
/// straight to the ALLOCATING binary_op (not the reuse-enabled add/sub/...): the `scalar(s)` temporary1426
/// is 0-d, so letting it bind a buffer-reuse overload would compute the result into the scalar and1427
/// collapse it to 0-d. The array operand here is a const lvalue (the caller keeps it), so the result1428
/// must be a fresh array of the BROADCAST shape regardless.1429
/**1430
* `a + s`: add scalar @p s to every element.1431
* @param a the array.1432
* @param s the arithmetic scalar (converted to @p T).1433
* @return a fresh array of @p a's shape.1434
* @complexity O(size of @p a).1435
* @alloc allocates the result and a one-element scalar temporary.1436
*/1437
template <typename T, typename S>1438
requires std::is_arithmetic_v<S>1439
basic_ndarray<T> operator+(const basic_ndarray<T>& a, S s) { return binary_op(a, scalar(static_cast<T>(s)), detail::add_op{}); }1440
/**1441
* `s + a`: add scalar @p s to every element.1442
* @param s the arithmetic scalar (converted to @p T).1443
* @param a the array.1444
* @return a fresh array of @p a's shape.1445
* @complexity O(size of @p a).1446
* @alloc allocates the result and a one-element scalar temporary.1447
*/1448
template <typename T, typename S>1449
requires std::is_arithmetic_v<S>1450
basic_ndarray<T> operator+(S s, const basic_ndarray<T>& a) { return binary_op(scalar(static_cast<T>(s)), a, detail::add_op{}); }1451
/**1452
* `a - s`: subtract scalar @p s from every element.1453
* @param a the array.1454
* @param s the arithmetic scalar (converted to @p T).1455
* @return a fresh array of @p a's shape.1456
* @complexity O(size of @p a).1457
* @alloc allocates the result and a one-element scalar temporary.1458
*/1459
template <typename T, typename S>1460
requires std::is_arithmetic_v<S>1461
basic_ndarray<T> operator-(const basic_ndarray<T>& a, S s) { return binary_op(a, scalar(static_cast<T>(s)), detail::sub_op{}); }1462
/**1463
* `s - a`: elementwise @p s minus each element.1464
* @param s the arithmetic scalar (converted to @p T).1465
* @param a the array.1466
* @return a fresh array of @p a's shape.1467
* @complexity O(size of @p a).1468
* @alloc allocates the result and a one-element scalar temporary.1469
* @test CheatahNDArray.ScalarTimesSizeOneArrayKeepsShape1470
*/1471
template <typename T, typename S>1472
requires std::is_arithmetic_v<S>1473
basic_ndarray<T> operator-(S s, const basic_ndarray<T>& a) { return binary_op(scalar(static_cast<T>(s)), a, detail::sub_op{}); }1474
/**1475
* `a * s`: multiply every element by scalar @p s.1476
* @param a the array.1477
* @param s the arithmetic scalar (converted to @p T).1478
* @return a fresh array of @p a's shape.1479
* @complexity O(size of @p a).1480
* @alloc allocates the result and a one-element scalar temporary.1481
* @test CheatahNDArray.CompoundAssignInPlace1482
* @crtest LangFeatures.NdarrayOperators1483
*/1484
template <typename T, typename S>1485
requires std::is_arithmetic_v<S>1486
basic_ndarray<T> operator*(const basic_ndarray<T>& a, S s) { return binary_op(a, scalar(static_cast<T>(s)), detail::mul_op{}); }1487
/**1488
* `s * a`: multiply every element by scalar @p s.1489
* @param s the arithmetic scalar (converted to @p T).1490
* @param a the array.1491
* @return a fresh array of @p a's shape.1492
* @complexity O(size of @p a).1493
* @alloc allocates the result and a one-element scalar temporary.1494
* @test CheatahNDArray.ScalarTimesSizeOneArrayKeepsShape1495
* @crtest LangFeatures.NdarrayOperators1496
*/1497
template <typename T, typename S>1498
requires std::is_arithmetic_v<S>1499
basic_ndarray<T> operator*(S s, const basic_ndarray<T>& a) { return binary_op(scalar(static_cast<T>(s)), a, detail::mul_op{}); }1500
/**1501
* `a / s`: divide every element by scalar @p s.1502
* @param a the array.1503
* @param s the arithmetic scalar divisor (converted to @p T).1504
* @return a fresh array of @p a's shape.1505
* @complexity O(size of @p a).1506
* @alloc allocates the result and a one-element scalar temporary.1507
*/1508
template <typename T, typename S>1509
requires std::is_arithmetic_v<S>1510
basic_ndarray<T> operator/(const basic_ndarray<T>& a, S s) { return binary_op(a, scalar(static_cast<T>(s)), detail::div_op{}); }1511
/**1512
* `s / a`: elementwise @p s divided by each element.1513
* @param s the arithmetic scalar numerator (converted to @p T).1514
* @param a the array of divisors.1515
* @return a fresh array of @p a's shape.1516
* @complexity O(size of @p a).1517
* @alloc allocates the result and a one-element scalar temporary.1518
*/1519
template <typename T, typename S>1520
requires std::is_arithmetic_v<S>1521
basic_ndarray<T> operator/(S s, const basic_ndarray<T>& a) { return binary_op(scalar(static_cast<T>(s)), a, detail::div_op{}); }1523
/// @cond INTERNAL1524
/// implementation plumbing / compiler-selected reuse overloads (README documents the public forms)1525
/// rvalue-reuse scalar forms: a temporary array operand is computed into IN PLACE (no alloc). Only the1526
/// commutative `s + a` / `s * a` get a scalar-LEFT reuse form; `s - a` / `s / a` keep the allocating1527
/// const& form (a reversed in-place would be needed, not worth it for that rare case).1528
/** `a + s` reusing the expiring array @p a in place. @param a expiring array operand. @param s the arithmetic scalar (converted to @p T). @return the result, in @p a's buffer. @complexity O(size of @p a). @alloc none on the fast path. */1529
template <typename T, typename S>1530
requires std::is_arithmetic_v<S>1531
basic_ndarray<T> operator+(basic_ndarray<T>&& a, S s) { return add(std::move(a), scalar(static_cast<T>(s))); }1532
/** `s + a` reusing the expiring array @p a in place (commutative). @param s the arithmetic scalar (converted to @p T). @param a expiring array operand. @return the result, in @p a's buffer. @complexity O(size of @p a). @alloc none on the fast path. */1533
template <typename T, typename S>1534
requires std::is_arithmetic_v<S>1535
basic_ndarray<T> operator+(S s, basic_ndarray<T>&& a) { return add(std::move(a), scalar(static_cast<T>(s))); }1536
/** `a - s` reusing the expiring array @p a in place. @param a expiring array operand. @param s the arithmetic scalar (converted to @p T). @return the result, in @p a's buffer. @complexity O(size of @p a). @alloc none on the fast path. */1537
template <typename T, typename S>1538
requires std::is_arithmetic_v<S>1539
basic_ndarray<T> operator-(basic_ndarray<T>&& a, S s) { return sub(std::move(a), scalar(static_cast<T>(s))); }1540
/** `a * s` reusing the expiring array @p a in place. @param a expiring array operand. @param s the arithmetic scalar (converted to @p T). @return the result, in @p a's buffer. @complexity O(size of @p a). @alloc none on the fast path. */1541
template <typename T, typename S>1542
requires std::is_arithmetic_v<S>1543
basic_ndarray<T> operator*(basic_ndarray<T>&& a, S s) { return mul(std::move(a), scalar(static_cast<T>(s))); }1544
/** `s * a` reusing the expiring array @p a in place (commutative). @param s the arithmetic scalar (converted to @p T). @param a expiring array operand. @return the result, in @p a's buffer. @complexity O(size of @p a). @alloc none on the fast path. */1545
template <typename T, typename S>1546
requires std::is_arithmetic_v<S>1547
basic_ndarray<T> operator*(S s, basic_ndarray<T>&& a) { return mul(std::move(a), scalar(static_cast<T>(s))); }1548
/** `a / s` reusing the expiring array @p a in place. @param a expiring array operand. @param s the arithmetic scalar divisor (converted to @p T). @return the result, in @p a's buffer. @complexity O(size of @p a). @alloc none on the fast path. */1549
template <typename T, typename S>1550
requires std::is_arithmetic_v<S>1551
basic_ndarray<T> operator/(basic_ndarray<T>&& a, S s) { return divide(std::move(a), scalar(static_cast<T>(s))); }1552
/// @endcond1554
/// In-place compound assignment: `a += b`, `a -= b`, `a *= b`, `a /= b`1555
/// (array or arithmetic-scalar right operand). See compound_apply.1556
/**1557
* In-place `a += b`, updating @p a's buffer (see compound_apply).1558
* @param a the array to update in place.1559
* @param b the right operand (same shape or single-element).1560
* @return reference to @p a.1561
* @complexity O(size of @p a).1562
* @alloc none on the contiguous fast path; the broadcast fallback allocates the result.1563
* @test CheatahNDArray.CompoundAssignInPlace1564
* @test CheatahNDArray.CompoundAssignNonContiguousFallback1565
* @crtest LangFeatures.NdarrayOperators1566
*/1567
template <typename T>1568
basic_ndarray<T>& operator+=(basic_ndarray<T>& a, const basic_ndarray<T>& b) {1569
compound_apply(a, b, detail::add_op{}); return a;1570
}1571
/**1572
* In-place `a -= b`, updating @p a's buffer (see compound_apply).1573
* @param a the array to update in place.1574
* @param b the right operand (same shape or single-element).1575
* @return reference to @p a.1576
* @complexity O(size of @p a).1577
* @alloc none on the contiguous fast path; the broadcast fallback allocates the result.1578
* @test CheatahNDArray.CompoundAssignNonContiguousFallback1579
*/1580
template <typename T>1581
basic_ndarray<T>& operator-=(basic_ndarray<T>& a, const basic_ndarray<T>& b) {1582
compound_apply(a, b, detail::sub_op{}); return a;1583
}1584
/**1585
* In-place `a *= b`, updating @p a's buffer (see compound_apply).1586
* @param a the array to update in place.1587
* @param b the right operand (same shape or single-element).1588
* @return reference to @p a.1589
* @complexity O(size of @p a).1590
* @alloc none on the contiguous fast path; the broadcast fallback allocates the result.1591
* @test CheatahNDArray.CompoundAssignNonContiguousFallback1592
*/1593
template <typename T>1594
basic_ndarray<T>& operator*=(basic_ndarray<T>& a, const basic_ndarray<T>& b) {1595
compound_apply(a, b, detail::mul_op{}); return a;1596
}1597
/**1598
* In-place `a /= b`, updating @p a's buffer (see compound_apply).1599
* @param a the array to update in place.1600
* @param b the right operand (same shape or single-element).1601
* @return reference to @p a.1602
* @complexity O(size of @p a).1603
* @alloc none on the contiguous fast path; the broadcast fallback allocates the result.1604
* @test CheatahNDArray.CompoundAssignNonContiguousFallback1605
*/1606
template <typename T>1607
basic_ndarray<T>& operator/=(basic_ndarray<T>& a, const basic_ndarray<T>& b) {1608
compound_apply(a, b, detail::div_op{}); return a;1609
}1610
/**1611
* In-place `a += s` (scalar).1612
* @param a the array to update.1613
* @param s the arithmetic scalar (converted to @p T).1614
* @return reference to @p a.1615
* @complexity O(size of @p a).1616
* @alloc allocates a one-element scalar temporary; the non-contiguous fallback also allocates the result.1617
*/1618
template <typename T, typename S>1619
requires std::is_arithmetic_v<S>1620
basic_ndarray<T>& operator+=(basic_ndarray<T>& a, S s) { return a += scalar(static_cast<T>(s)); }1621
/**1622
* In-place `a -= s` (scalar).1623
* @param a the array to update.1624
* @param s the arithmetic scalar (converted to @p T).1625
* @return reference to @p a.1626
* @complexity O(size of @p a).1627
* @alloc allocates a one-element scalar temporary; the non-contiguous fallback also allocates the result.1628
* @test CheatahNDArray.CompoundAssignInPlace1629
*/1630
template <typename T, typename S>1631
requires std::is_arithmetic_v<S>1632
basic_ndarray<T>& operator-=(basic_ndarray<T>& a, S s) { return a -= scalar(static_cast<T>(s)); }1633
/**1634
* In-place `a *= s` (scalar).1635
* @param a the array to update.1636
* @param s the arithmetic scalar (converted to @p T).1637
* @return reference to @p a.1638
* @complexity O(size of @p a).1639
* @alloc allocates a one-element scalar temporary; the non-contiguous fallback also allocates the result.1640
* @test CheatahNDArray.CompoundAssignInPlace1641
* @crtest LangFeatures.ParamsPassByReference1642
*/1643
template <typename T, typename S>1644
requires std::is_arithmetic_v<S>1645
basic_ndarray<T>& operator*=(basic_ndarray<T>& a, S s) { return a *= scalar(static_cast<T>(s)); }1646
/**1647
* In-place `a /= s` (scalar).1648
* @param a the array to update.1649
* @param s the arithmetic scalar divisor (converted to @p T).1650
* @return reference to @p a.1651
* @complexity O(size of @p a).1652
* @alloc allocates a one-element scalar temporary; the non-contiguous fallback also allocates the result.1653
* @test CheatahNDArray.CompoundAssignInPlace1654
* @crtest LangFeatures.NdarrayOperators1655
*/1656
template <typename T, typename S>1657
requires std::is_arithmetic_v<S>1658
basic_ndarray<T>& operator/=(basic_ndarray<T>& a, S s) { return a /= scalar(static_cast<T>(s)); }1660
// ---- complex support ----1661
namespace detail {1662
/// Map @p a element-wise through @p f into a fresh contiguous array of element type1663
/// `U` (which may differ from `T` — e.g. complex→real for @ref real). Contiguous1664
/// fast path via `std::transform(unseq)`; otherwise a C-order walk.1665
template <typename U, Field T, typename F>1666
basic_ndarray<U> map_array(const basic_ndarray<T>& a, F f) {1667
basic_ndarray<U> out = basic_ndarray<U>::uninitialized(a.shape()); // fully written below1668
auto& obuf = *out.buffer();1669
if (is_contiguous(a)) {1670
const auto& abuf = *a.buffer();1671
std::transform(CHEATAH_UNSEQ abuf.begin() + a.offset(),1672
abuf.begin() + a.offset() + a.size(), obuf.begin(), f);1673
return out;1674
}1675
std::vector<std::size_t> idx(a.ndim(), 0);1676
std::size_t flat = 0;1677
for (;;) { // a 0-d array still yields its one element1678
obuf[flat++] = f(a.at(idx));1679
if (a.ndim() == 0 || !next_index(idx, a.shape())) break;1680
}1681
return out;1682
}1684
// Out-of-line, separately-compiled (-ffast-math) double-precision SIMD kernels for the1685
// element-wise ufuncs — see ufunc_simd.cpp. They vectorize the transcendentals through1686
// libmvec, which the default flags cannot; isolating -ffast-math to that file keeps the1687
// rest of cheatah's arithmetic strict.1688
void simd_sqrt_f64(const double*, double*, std::size_t);1689
void simd_cbrt_f64(const double*, double*, std::size_t);1690
void simd_exp_f64(const double*, double*, std::size_t);1691
void simd_log_f64(const double*, double*, std::size_t);1692
void simd_sin_f64(const double*, double*, std::size_t);1693
void simd_cos_f64(const double*, double*, std::size_t);1694
void simd_tan_f64(const double*, double*, std::size_t);1696
/// Map a ufunc over @p a: a *contiguous double* array goes through the precompiled SIMD1697
/// @p kernel; everything else (float, or a strided/broadcast view) uses the generic1698
/// scalar @p fallback. Same result either way — the kernel just vectorizes the hot case.1699
template <FloatingPoint T, class Kernel, class Fallback>1700
basic_ndarray<T> map_ufunc(const basic_ndarray<T>& a, Kernel kernel, Fallback fallback) {1701
if constexpr (std::is_same_v<T, double>) {1702
if (is_contiguous(a)) {1703
basic_ndarray<T> out = basic_ndarray<T>::uninitialized(a.shape()); // kernel fills it1704
kernel(a.buffer()->data() + a.offset(), out.buffer()->data(), a.size());1705
return out;1706
}1707
}1708
return map_array<T>(a, fallback);1709
}1710
} // namespace detail1712
/**1713
* Build a complex array from real and imaginary parts (element-wise `re + im·j`),1714
* broadcasting the two together — the way to construct a complex matrix/vector1715
* (a wavefunction, a Hermitian operator) since cheatah literals are real.1716
* @param re the real parts (a real floating array).1717
* @param im the imaginary parts (a real floating array, broadcast against @p re).1718
* @return a `basic_ndarray<std::complex<T>>` of `re + im·j`; throws if the shapes don't broadcast.1719
* @complexity O(size of result).1720
* @alloc allocates the result buffer.1721
* @test CheatahNDArray.ComplexConstructAndParts1722
* @crtest NdarrayCompileRun.Complex1723
* @systest StdlibE2E.NdarrayComplex1724
*/1725
template <FloatingPoint T>1726
basic_ndarray<std::complex<T>> complex(const basic_ndarray<T>& re, const basic_ndarray<T>& im) {1727
using C = std::complex<T>;1728
const std::vector<std::size_t> rshape = broadcast_shapes(re.shape(), im.shape());1729
const basic_ndarray<T> rv = broadcast_to(re, rshape);1730
const basic_ndarray<T> iv = broadcast_to(im, rshape);1731
basic_ndarray<C> out(rshape);1732
auto& obuf = *out.buffer();1733
std::vector<std::size_t> idx(rshape.size(), 0);1734
std::size_t flat = 0;1735
for (;;) { // a 0-d result still has its one element1736
obuf[flat++] = C(rv.at(idx), iv.at(idx));1737
if (rshape.empty() || !detail::next_index(idx, rshape)) break;1738
}1739
return out;1740
}1741
/**1742
* Element-wise complex conjugate (`a − b·j` for each `a + b·j`); on a real array it1743
* is the identity (a copy). Type-preserving. Used to form Hermitian adjoints and1744
* conjugate-linear inner products.1745
* @param a the array.1746
* @return a fresh array of the same element type with each element conjugated.1747
* @complexity O(size).1748
* @alloc allocates the result buffer.1749
* @test CheatahNDArray.ComplexConstructAndParts1750
* @crtest NdarrayCompileRun.Conj1751
* @systest StdlibE2E.NdarrayComplex1752
*/1753
template <Field T>1754
basic_ndarray<T> conj(const basic_ndarray<T>& a) {1755
return detail::map_array<T>(a, [](T x) -> T {1756
if constexpr (is_complex_v<T>) {1757
return std::conj(x);1758
} else {1759
return x;1760
}1761
});1762
}1763
/**1764
* The real parts as a real array (the identity on a real array). For `a + b·j` it1765
* returns `a`.1766
* @param a the array.1767
* @return a `basic_ndarray<real_base_t<T>>` of the real parts.1768
* @complexity O(size).1769
* @alloc allocates the result buffer.1770
* @test CheatahNDArray.ComplexConstructAndParts1771
* @crtest NdarrayCompileRun.Real1772
* @systest StdlibE2E.NdarrayComplex1773
*/1774
template <Field T>1775
basic_ndarray<real_base_t<T>> real(const basic_ndarray<T>& a) {1776
using R = real_base_t<T>;1777
return detail::map_array<R>(a, [](T x) -> R {1778
if constexpr (is_complex_v<T>) {1779
return x.real();1780
} else {1781
return x;1782
}1783
});1784
}1785
/**1786
* The imaginary parts as a real array (all zeros for a real array). For `a + b·j` it1787
* returns `b`.1788
* @param a the array.1789
* @return a `basic_ndarray<real_base_t<T>>` of the imaginary parts.1790
* @complexity O(size).1791
* @alloc allocates the result buffer.1792
* @test CheatahNDArray.ComplexConstructAndParts1793
* @crtest NdarrayCompileRun.Imag1794
* @systest StdlibE2E.NdarrayComplex1795
*/1796
template <Field T>1797
basic_ndarray<real_base_t<T>> imag(const basic_ndarray<T>& a) {1798
using R = real_base_t<T>;1799
return detail::map_array<R>(a, [](T x) -> R {1800
if constexpr (is_complex_v<T>) {1801
return x.imag();1802
} else {1803
return R{0};1804
}1805
});1806
}1808
// ---- element-wise math (numpy-style ufuncs) ----1809
// These are the array counterparts of the scalar `math` module — mirroring Python's1810
// split: `math.sqrt(x)` for a scalar, `ndarray.sqrt(a)` (≈ `numpy.sqrt`) for a whole1811
// array. A contiguous `double` array routes through a precompiled SIMD kernel1812
// (ufunc_simd.cpp) that vectorizes via glibc's libmvec — so `exp`/`sin`/… run at vector1813
// speed and beat NumPy's ufuncs; other element types / strided views fall back to a1814
// scalar map (see detail::map_ufunc).1815
/**1816
* Element-wise square root (the array form of `math.sqrt`; ≈ `numpy.sqrt`).1817
* @param a a floating-point array.1818
* @return a fresh same-shape array with `√x` for each element.1819
* @complexity O(size). @alloc allocates the result buffer.1820
* @test CheatahNDArray.ElementwiseMath1821
* @crtest NdarrayCompileRun.Sqrt1823
*/1824
template <FloatingPoint T>1825
basic_ndarray<T> sqrt(const basic_ndarray<T>& a) {1826
return detail::map_ufunc<T>(a, detail::simd_sqrt_f64, [](T x) { return std::sqrt(x); });1827
}1828
/**1829
* Element-wise cube root (the array form of `math.cbrt`; ≈ `numpy.cbrt`).1830
* @param a a floating-point array.1831
* @return a fresh same-shape array with `∛x` for each element.1832
* @complexity O(size). @alloc allocates the result buffer.1833
* @test CheatahNDArray.ElementwiseMath1835
*/1836
template <FloatingPoint T>1837
basic_ndarray<T> cbrt(const basic_ndarray<T>& a) {1838
return detail::map_ufunc<T>(a, detail::simd_cbrt_f64, [](T x) { return std::cbrt(x); });1839
}1840
/**1841
* Element-wise eˣ (the array form of `math.exp`; ≈ `numpy.exp`).1842
* @param a a floating-point array.1843
* @return a fresh same-shape array with `exp(x)` for each element.1844
* @complexity O(size). @alloc allocates the result buffer.1845
* @test CheatahNDArray.ElementwiseMath1846
* @crtest NdarrayCompileRun.Exp1848
*/1849
template <FloatingPoint T>1850
basic_ndarray<T> exp(const basic_ndarray<T>& a) {1851
return detail::map_ufunc<T>(a, detail::simd_exp_f64, [](T x) { return std::exp(x); });1852
}1853
/**1854
* Element-wise natural log (the array form of `math.log`; ≈ `numpy.log`).1855
* @param a a floating-point array.1856
* @return a fresh same-shape array with `ln(x)` for each element.1857
* @complexity O(size). @alloc allocates the result buffer.1858
* @test CheatahNDArray.ElementwiseMath1859
*/1860
template <FloatingPoint T>1861
basic_ndarray<T> log(const basic_ndarray<T>& a) {1862
return detail::map_ufunc<T>(a, detail::simd_log_f64, [](T x) { return std::log(x); });1863
}1864
/**1865
* Element-wise sine (the array form of `math.sin`; ≈ `numpy.sin`).1866
* @param a a floating-point array (radians).1867
* @return a fresh same-shape array with `sin(x)` for each element.1868
* @complexity O(size). @alloc allocates the result buffer.1869
* @test CheatahNDArray.ElementwiseMath1870
* @crtest NdarrayCompileRun.Sin1871
* @systest StdlibE2E.NdarrayMath1872
*/1873
template <FloatingPoint T>1874
basic_ndarray<T> sin(const basic_ndarray<T>& a) {1875
return detail::map_ufunc<T>(a, detail::simd_sin_f64, [](T x) { return std::sin(x); });1876
}1877
/**1878
* Element-wise cosine (the array form of `math.cos`; ≈ `numpy.cos`).1879
* @param a a floating-point array (radians).1880
* @return a fresh same-shape array with `cos(x)` for each element.1881
* @complexity O(size). @alloc allocates the result buffer.1882
* @test CheatahNDArray.ElementwiseMath1883
*/1884
template <FloatingPoint T>1885
basic_ndarray<T> cos(const basic_ndarray<T>& a) {1886
return detail::map_ufunc<T>(a, detail::simd_cos_f64, [](T x) { return std::cos(x); });1887
}1888
/**1889
* Element-wise tangent (the array form of `math.tan`; ≈ `numpy.tan`).1890
* @param a a floating-point array (radians).1891
* @return a fresh same-shape array with `tan(x)` for each element.1892
* @complexity O(size). @alloc allocates the result buffer.1893
* @test CheatahNDArray.ElementwiseMath1894
*/1895
template <FloatingPoint T>1896
basic_ndarray<T> tan(const basic_ndarray<T>& a) {1897
return detail::map_ufunc<T>(a, detail::simd_tan_f64, [](T x) { return std::tan(x); });1898
}1899
/**1900
* Element-wise absolute value (the array form of `math.abs`; ≈ `numpy.abs`).1901
* @param a a floating-point array.1902
* @return a fresh same-shape array with `|x|` for each element.1903
* @complexity O(size). @alloc allocates the result buffer.1904
* @test CheatahNDArray.ElementwiseMath1905
* @systest StdlibE2E.NdarrayMath1906
*/1907
template <FloatingPoint T>1908
basic_ndarray<T> abs(const basic_ndarray<T>& a) {1909
return detail::map_array<T>(a, [](T x) { return std::fabs(x); });1910
}1912
// ---- reductions / access / display ----1913
namespace detail {1914
/// The shared multi-accumulator reduction: sums `get(0)..get(n-1)` with EIGHT independent1915
/// accumulators, tree-combined, plus a scalar tail. The independent lanes break the FP-add1916
/// dependency chain so -O3 -march=native emits SIMD+FMA and reaches memory bandwidth instead of1917
/// add latency (a single running sum — or a plain `std::reduce`, which libstdc++ left-folds for FP1918
/// without -ffast-math — serializes: the dot/norm mistake). `get(i)` returns the i-th TERM — an1919
/// element for `sum`, a (possibly conjugated) product for `dot`, a strided read for `trace`. One1920
/// primitive replaces the copies formerly hand-rolled in ndarray/linalg. `constexpr`, so a1921
/// fixed-extent caller gets a compile-time reduction too.1922
template <class T, class Get>1923
constexpr T reduce_lanes(std::size_t n, Get get) {1924
T s0{}, s1{}, s2{}, s3{}, s4{}, s5{}, s6{}, s7{};1925
std::size_t i = 0;1926
for (; i + 8 <= n; i += 8) {1927
s0 += get(i + 0); s1 += get(i + 1); s2 += get(i + 2); s3 += get(i + 3);1928
s4 += get(i + 4); s5 += get(i + 5); s6 += get(i + 6); s7 += get(i + 7);1929
}1930
T s = ((s0 + s1) + (s2 + s3)) + ((s4 + s5) + (s6 + s7));1931
for (; i < n; ++i) s += get(i);1932
return s;1933
}1934
} // namespace detail1935
/**1936
* Sum of all elements — a full reduction across every axis (a contiguous array goes1937
* through the shared multi-accumulator SIMD reduction @ref detail::reduce_lanes, else1938
* a C-order walk); empty sums to 0.1939
* @param a the array.1940
* @return the total, as the element type @p T.1941
* @complexity O(size).1942
* @alloc none on the contiguous fast path; the strided walk allocates an index vector.1943
* @test CheatahNDArray.ShapeFactoriesAndReductions1944
* @crtest NdarrayCompileRun.Sum1945
* @systest StdlibE2E.Ndarray1946
*/1947
template <Field T>1948
T sum(const basic_ndarray<T>& a) {1949
if (is_contiguous(a)) {1950
// Contiguous fast path via the shared multi-accumulator reduction (each term is one element).1951
const T* p = a.buffer()->data() + a.offset();1952
return detail::reduce_lanes<T>(a.size(), [p](std::size_t i) { return p[i]; });1953
}1954
T s{};1955
std::vector<std::size_t> idx(a.ndim(), 0);1956
for (;;) { // a 0-d array still contributes its one element1957
s += a.at(idx);1958
if (a.ndim() == 0 || !detail::next_index(idx, a.shape())) break;1959
}1960
return s;1961
}1962
/**1963
* Mean of all elements, always as a `double` (0.0 for an empty array — no divide-by-zero).1964
* @param a the array.1965
* @return the average as a double.1966
* @complexity O(size).1967
* @alloc none on a contiguous array; a strided source allocates sum's index vector.1968
* @test CheatahNDArray.ShapeFactoriesAndReductions1969
* @crtest NdarrayCompileRun.Mean1970
* @systest StdlibE2E.Ndarray1971
*/1972
template <Numeric T>1973
double mean(const basic_ndarray<T>& a) {1974
const std::size_t n = a.size();1975
return n == 0 ? 0.0 : static_cast<double>(sum(a)) / static_cast<double>(n);1976
}1977
/**1978
* Read one element by signed multi-index (the cheatah-facing wrapper over @ref1979
* basic_ndarray::at; rejects negative coordinates).1980
* @param a the array.1981
* @param index one coordinate per dimension (signed; throws on a negative).1982
* @return the element value (type @p T); throws on a wrong-rank/out-of-range index.1983
* @complexity O(ndim).1984
* @alloc none.1985
* @test CheatahNDArray.ShapeFactoriesAndReductions1986
* @crtest NdarrayCompileRun.Get1987
* @systest StdlibE2E.Ndarray1988
*/1989
template <Copyable T>1990
T get(const basic_ndarray<T>& a, const std::vector<long long>& index) {1991
return a.at(detail::to_size(index));1992
}1993
/**1994
* The shape as signed dims (cheatah integers are signed; a 0-d array yields an empty list).1995
* @param a the array.1996
* @return the dimensions as a `long long` vector.1997
* @complexity O(ndim).1998
* @alloc allocates the result vector.1999
* @test CheatahNDArray.ShapeFactoriesAndReductions2000
* @crtest NdarrayCompileRun.ShapeOf2001
* @systest StdlibE2E.Ndarray2002
*/2003
template <Element T>2004
std::vector<long long> shape_of(const basic_ndarray<T>& a) {2005
std::vector<long long> out(a.ndim());2006
for (std::size_t i = 0; i < a.ndim(); ++i) out[i] = static_cast<long long>(a.shape()[i]);2007
return out;2008
}2009
/**2010
* The element count as a signed value (1 for a 0-d array).2011
* @param a the array.2012
* @return the number of elements as a `long long`.2013
* @complexity O(ndim).2014
* @alloc none.2015
* @test CheatahNDArray.ShapeFactoriesAndReductions2016
* @crtest NdarrayCompileRun.SizeOf2017
* @crtest NdarrayCompileRun.RprintShowsArrayFull2018
*/2019
template <Element T>2020
long long size_of(const basic_ndarray<T>& a) {2021
return static_cast<long long>(a.size());2022
}2024
namespace detail {2025
/// Format one element. Real types go through `operator<<`; a complex element is2026
/// rendered Python-style as `a+bj` / `a-bj` (not the `std::complex` default2027
/// `(a,b)`), so a complex spectrum prints the way a cheatah user expects.2028
template <typename T>2029
std::string format_scalar(const T& v) {2030
std::ostringstream os;2031
if constexpr (is_complex_v<T>) {2032
using R = real_base_t<T>;2033
// Flush negative zero to +0 so a conjugate prints "1+0j", not "1+-0j".2034
const auto nz = [](R x) -> R { return x == R{0} ? R{0} : x; };2035
os << nz(v.real());2036
if (v.imag() < R{0}) {2037
os << "-" << nz(-v.imag()) << "j";2038
} else {2039
os << "+" << nz(v.imag()) << "j";2040
}2041
} else if constexpr (std::is_same_v<T, signed char> || std::is_same_v<T, unsigned char> ||2042
std::is_same_v<T, char>) {2043
os << +v; // i8/u8 are char-sized: promote so an element prints as a NUMBER, not a character2044
} else {2045
os << v;2046
}2047
return os.str();2048
}2050
/// Recursively format @p a into nested brackets (each element via `format_scalar`).2051
template <Element T>2052
void format_rec(const basic_ndarray<T>& a, std::vector<std::size_t>& idx, std::size_t dim,2053
std::string& out) {2054
if (dim == a.ndim()) {2055
out += format_scalar(a.at(idx));2056
return;2057
}2058
out += "[";2059
for (std::size_t i = 0; i < a.shape()[dim]; ++i) {2060
if (i != 0) out += ", ";2061
idx[dim] = i;2062
format_rec(a, idx, dim + 1, out);2063
}2064
out += "]";2065
}2067
/// Like format_rec but ABBREVIATES a large array: an axis longer than `2*edge` shows its2068
/// first and last `edge` items with `...` between, recursively. Summarization is enabled by2069
/// @p summarize (the caller turns it on only past a total-size threshold), so small arrays2070
/// print in full.2071
template <Element T>2072
void format_rec_trunc(const basic_ndarray<T>& a, std::vector<std::size_t>& idx, std::size_t dim,2073
std::string& out, std::size_t edge, bool summarize) {2074
if (dim == a.ndim()) {2075
out += format_scalar(a.at(idx));2076
return;2077
}2078
const std::size_t n = a.shape()[dim];2079
const bool trunc = summarize && n > 2 * edge;2080
out += "[";2081
bool first = true;2082
for (std::size_t i = 0; i < n; ++i) {2083
if (trunc && i >= edge && i < n - edge) {2084
if (i == edge) {2085
if (!first) out += ", ";2086
out += "...";2087
first = false;2088
}2089
continue; // skip the abbreviated middle2090
}2091
if (!first) out += ", ";2092
first = false;2093
idx[dim] = i;2094
format_rec_trunc(a, idx, dim + 1, out, edge, summarize);2095
}2096
out += "]";2097
}2099
/// to_string, but ABBREVIATED with `...` when the array is large (total size beyond a2100
/// numpy-style threshold) — the readable default for `io.print`. `io.rprint`/`to_string`2101
/// keep the full form.2102
template <Element T>2103
std::string to_string_pretty(const basic_ndarray<T>& a) {2104
if (a.ndim() == 0) return format_scalar(a.at({}));2105
constexpr std::size_t kEdge = 3; // items kept at each end of an abbreviated axis2106
constexpr std::size_t kThreshold = 1000; // only summarize past this many elements (numpy)2107
std::vector<std::size_t> idx(a.ndim(), 0);2108
std::string out;2109
format_rec_trunc(a, idx, 0, out, kEdge, a.size() > kThreshold);2110
return out;2111
}2112
} // namespace detail2114
/**2115
* Render as a nested-bracket string, e.g. `"[[1, 2], [3, 4]]"` (a 0-d scalar renders2116
* as the bare number). Each element is formatted with the default `ostream` precision.2117
* @param a the array.2118
* @return the textual representation.2119
* @complexity O(size).2120
* @alloc allocates the result string, a formatting stream per element, and an index vector.2121
* @test CheatahNDArray.ToStringScalar, CheatahNDArray.BroadcastingAdd2122
* @crtest NdarrayCompileRun.ToString2123
* @systest StdlibE2E.Ndarray2124
*/2125
template <Element T>2126
std::string to_string(const basic_ndarray<T>& a) {2127
if (a.ndim() == 0) {2128
return detail::format_scalar(a.at({}));2129
}2130
std::vector<std::size_t> idx(a.ndim(), 0);2131
std::string out;2132
detail::format_rec(a, idx, 0, out);2133
return out;2134
}2136
template <Element T>2137
inline std::string basic_ndarray<T>::str() const {2138
return to_string(*this);2139
}2141
/**2142
* Stream an array to a `std::ostream` (the FULL nested-bracket form) — so an NDArray is2143
* directly Streamable, like a primitive or a cheatah struct, without going through2144
* `to_string`/`str()`. `io.rprint`, `str()`, and a struct that holds an array all stream it2145
* this way; `io.print` instead uses @ref cheatah_pretty_print to abbreviate large arrays.2146
* @param os the stream.2147
* @param a the array.2148
* @return @p os.2149
* @complexity O(size).2150
* @alloc allocates the intermediate string.2151
* @test CheatahNDArray.StreamableOperator2152
* @systest StdlibE2E.Ndarray2153
*/2154
template <Element T>2155
std::ostream& operator<<(std::ostream& os, const basic_ndarray<T>& a) {2156
return os << to_string(a);2157
}2159
template <Element T>2160
inline void basic_ndarray<T>::cheatah_pretty_print(std::ostream& os, long long /*unused*/) const {2161
os << detail::to_string_pretty(*this);2162
}2164
} // namespace cheatah::ndarray2166
// cheatah's value-position subscript lowers to builtins::index(obj, i, ...).2167
// These overloads give it the ndarray meaning: negative-aware element reads,2168
// one coordinate per dimension.2169
namespace cheatah::builtins {2171
/// @cond INTERNAL2172
/// This header deliberately includes no project header — it only REOPENS the builtins namespace to2173
/// add ndarray overloads, exactly as `index` below does. So the two names the slice-assignment2174
/// vocabulary needs are spelled locally: `empty_seq` is forward-declared (a reference to an2175
/// incomplete type is all an overload declaration needs, and builtins.hpp completes it wherever2176
/// both headers are in play), and the "to the end" sentinel is the same value builtins::slice_end2177
/// holds. A static_assert in builtins.hpp keeps the two from drifting.2178
struct empty_seq;2179
inline constexpr long long nd_slice_end = std::numeric_limits<long long>::max();2180
/// A slice bound, with a negative counting back from the end (builtins::detail::norm_index).2181
inline long long nd_norm_index(long long i, long long n) { return i < 0 ? i + n : i; }2182
/// @endcond2185
/**2186
* Element read `a[i, j, ...]` (negative indices count from the dimension end).2187
* @param a the array to read from.2188
* @param first the first-axis coordinate (negative counts from that dimension's end).2189
* @param rest the remaining per-axis coordinates (one per further dimension).2190
* @return the element value.2191
* @complexity O(ndim). @alloc none.2192
* @test CheatahNDArray.SubscriptReadWrite2193
* @crtest LangFeatures.NdarraySubscript2194
* @systest StdlibE2E.Ndarray2195
*/2196
template <typename T, ::cheatah::ndarray::Subscript First, ::cheatah::ndarray::Subscript... Ix>2197
T index(const ::cheatah::ndarray::basic_ndarray<T>& a, First first, Ix... rest) {2198
// Any index may be a scoped enum column label — item_ref performs the one sanctioned conversion2199
// (ndarray::Subscript); the const overload reads, and the value is copied out.2200
return a.item_ref(first, rest...);2201
}2204
/**2205
* A view of `a[lo:hi]` along axis 0 — the rows `lo` up to `hi`, sharing @p a's buffer.2206
*2207
* Zero-copy, like every other ndarray view: only shape and offset change, the strides are @p a's,2208
* and the shared buffer keeps the elements alive for as long as either array refers to them.2209
* Writing through the view therefore writes into @p a. Bounds follow the list rules — negatives2210
* count from the end, out-of-range values clamp, a reversed range yields an empty leading axis.2211
* @tparam T the element type.2212
* @param a the array to view.2213
* @param lo first row (negative counts from the end).2214
* @param hi one past the last row, or `slice_end` for "to the end".2215
* @return a view sharing @p a's storage.2216
* @complexity O(ndim) — shape and stride vectors are copied, never the elements.2217
* @alloc the shape/stride vectors of the view; no element copy.2218
* @test CheatahNDArray.SliceIsAView2219
* @crtest LangFeatures.NdarraySliceAssignment2220
* @systest StdlibE2E.Ndarray2221
*/2222
template <typename T>2223
::cheatah::ndarray::basic_ndarray<T> slice(const ::cheatah::ndarray::basic_ndarray<T>& a,2224
long long lo, long long hi) {2225
if (a.ndim() == 0) throw std::runtime_error("ndarray: cannot slice a 0-d array");2226
const auto n = static_cast<long long>(a.shape()[0]);2227
lo = nd_norm_index(lo, n);2228
hi = (hi == nd_slice_end) ? n : nd_norm_index(hi, n);2229
if (lo < 0) lo = 0;2230
if (lo > n) lo = n;2231
if (hi > n) hi = n;2232
if (hi < lo) hi = lo;2233
std::vector<std::size_t> shape = a.shape();2234
shape[0] = static_cast<std::size_t>(hi - lo);2235
const std::size_t off = a.offset() + static_cast<std::size_t>(lo * a.strides()[0]);2236
return ::cheatah::ndarray::basic_ndarray<T>(a.buffer(), shape, a.strides(), off);2237
}2239
/**2240
* Write @p rhs into the elements `a[lo:hi]` addresses — an array assignment COPIES, it never2241
* rebinds or resizes.2242
*2243
* @p rhs may be a scalar (filling every addressed element) or an array broadcastable to the2244
* slice's shape, matching numpy. The destination keeps its shape, so an extent that does not2245
* broadcast is an error rather than a silent truncation. Strides are honoured, so writing2246
* through a non-contiguous view lands on the right elements.2247
* @tparam T the element type.2248
* @tparam R the source: a scalar or an ndarray.2249
* @param a the array to write into.2250
* @param lo first row (negative counts from the end).2251
* @param hi one past the last row, or `slice_end` for "to the end".2252
* @param rhs the value(s) to copy in.2253
* @complexity O(size of the slice).2254
* @alloc the view's shape/stride vectors, plus a broadcast view of @p rhs when it is an array.2255
* @test CheatahNDArray.SliceAssignCopiesIn2256
* @crtest LangFeatures.NdarraySliceAssignment2257
* @systest StdlibE2E.Ndarray2258
*/2259
template <typename T, typename R>2260
void slice_assign(::cheatah::ndarray::basic_ndarray<T>& a, long long lo, long long hi,2261
const R& rhs) {2262
::cheatah::ndarray::basic_ndarray<T> dst = slice(a, lo, hi);2263
if (dst.size() == 0) return;2264
std::vector<std::size_t> idx(dst.ndim(), 0);2265
if constexpr (std::is_convertible_v<R, T>) { // scalar fill2266
const T v = static_cast<T>(rhs);2267
for (;;) {2268
dst.at_ref(idx) = v;2269
if (!::cheatah::ndarray::detail::next_index(idx, dst.shape())) break;2270
}2271
} else { // array source, broadcast to the slice2272
const ::cheatah::ndarray::basic_ndarray<T> src =2273
::cheatah::ndarray::broadcast_to(rhs, dst.shape());2274
for (;;) {2275
dst.at_ref(idx) = src.at(idx);2276
if (!::cheatah::ndarray::detail::next_index(idx, dst.shape())) break;2277
}2278
}2279
}2281
/// @cond INTERNAL2282
/// An ndarray has a fixed shape, so there is nothing for `a[lo:hi] = []` to delete.2283
template <typename T>2284
void slice_assign(::cheatah::ndarray::basic_ndarray<T>& a, long long lo, long long hi,2285
const empty_seq& /*empty*/) {2286
(void)a; (void)lo; (void)hi;2287
throw std::runtime_error(2288
"ndarray: a[lo:hi] = [] has no meaning — an array assignment fills the slice, it cannot "2289
"remove elements (its shape is fixed)");2290
}2291
/// @endcond2293
} // namespace cheatah::builtins