cheatah
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 once
5/**
6 * @file ndarray.hpp
7 * @brief cheatah `ndarray` — our own numpy-flavored N-dimensional array
8 * (`basic_ndarray<T>` over any @ref Element type; `NDArray` is the `double`
9 * default) with NumPy broadcasting, surfaced as a `NDArray` class plus free
10 * 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 shared
18 * buffer (`std::shared_ptr<buffer_t<T>>`) and an array is a VIEW into
19 * 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, so
21 * 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 allocator
34#include <numeric>
35#include <sstream>
36#include <stdexcept>
37#include <string>
38#include <type_traits>
39#include <utility> // std::forward / std::move
40#include <version> // __cpp_lib_execution feature-test macro
41#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 the
45// feature-test macro and fall back to the plain (policy-less) overloads where it's
46// absent. This is speed-neutral: `unseq` is unsequenced (no threads, no TBB) and for
47// these simple element-wise loops the -O3 -march=native auto-vectorizer produces the
48// same SIMD either way — the transcendental vectorization comes from ufunc_simd.cpp's
49// libmvec/Accelerate kernels, not from this policy.
50#if defined(__cpp_lib_execution)
51#include <execution>
52#define CHEATAH_UNSEQ std::execution::unseq,
53#else
54#define CHEATAH_UNSEQ
55#endif
57namespace cheatah::ndarray {
59/// Numeric<T>: an arithmetic element type an ndarray can store (int or float
60/// family). Storage, construction, and elementwise +-* require only this.
61template <typename T>
62concept 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 an
65/// integer array fails with a clear "FloatingPoint not satisfied", not template spam.
66template <typename T>
67concept FloatingPoint = std::floating_point<T>;
69/// @cond INTERNAL
70template <typename T>
71struct is_complex : std::false_type {};
72template <typename U>
73struct is_complex<std::complex<U>> : std::bool_constant<std::is_floating_point_v<U>> {};
74/// @endcond
76/// Whether `T` is a `std::complex` of a floating type — the trait behind @ref Field.
77template <typename T>
78inline constexpr bool is_complex_v = is_complex<T>::value;
80/// Field<T>: a scalar an ndarray can store — a real arithmetic type OR a
81/// `std::complex` of a floating type. This is what makes **complex** matrices and
82/// vectors first-class (Hermitian operators, complex wavefunctions), and lets a
83/// REAL matrix yield the COMPLEX eigenvalues it mathematically has.
84template <typename T>
85concept 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/complex
88/// number) OR any MOVABLE type, so a fixed-size struct (a 2-D point, an RGBA colour, a GPU
89/// vertex) lives in an ndarray too. Elements are MOVED into the buffer on construction and the
90/// 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 constrained
92/// to @ref Field and the duplicating factories to @ref Copyable, so a move-only element still
93/// stores / indexes / views / moves — it simply cannot be summed or deep-copied, and the compiler
94/// says so by design (cheatah discourages hidden copies on hot data).
95template <typename T>
96concept Element = Field<T> || std::movable<T>;
98/// Copyable<T>: an @ref Element that may ALSO be duplicated. It gates only the value-fill
99/// 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 an
101/// 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).
103template <typename T>
104concept Copyable = Element<T> && std::copyable<T>;
106/// Subscript<T>: what may address an axis — an integer, OR a scoped `enum class` whose ordinal names
107/// the position (a column label). This concept is the ONLY door through which a scoped enum becomes an
108/// integer: `enum class` values stay strongly typed everywhere else, and the implicit
109/// enum-to-index conversion is confined to array subscripting, exactly where a named column belongs.
110template <typename T>
111concept 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.EnumIndexingOnVectorsAndMatrices
120template <Subscript Ix>
121[[nodiscard]] constexpr long long subscript_index(Ix i) noexcept {
122 return static_cast<long long>(i);
125/// @cond INTERNAL
126template <typename T>
127struct real_base {
128 using type = T;
129};
130template <typename U>
131struct real_base<std::complex<U>> {
132 using type = U;
133};
134/// @endcond
136/// The real type underlying a @ref Field `T` (`double` for both `double` and
137/// `complex<double>`).
138template <typename T>
139using real_base_t = typename real_base<T>::type;
141/// complex_of_t<T>: the complex type over T's real base. `eig`/`eigvals` return an
142/// array of these, because a real matrix can have complex eigenvalues (conjugate
143/// pairs) — e.g. the rotation matrix [[0,-1],[1,0]] has eigenvalues ±i.
144template <typename T>
145using complex_of_t = std::complex<real_base_t<T>>;
147namespace detail {
148/// @cond INTERNAL
149/// An allocator identical to `std::allocator<T>` in every respect EXCEPT that
150/// DEFAULT (no-value) construction — what `vector(n)` / `resize(n)` perform — leaves a
151/// 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 (binary
154/// ops, ufuncs, reshape, array()) would otherwise pay a throwaway zero-fill of the whole
155/// buffer first. That wasted write pass is hidden on compute-heavy ops but DOMINATES
156/// bandwidth-bound ones — `add` was ≈1.5× of NumPy purely from the extra memset. With
157/// 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.
159template <typename T>
160struct 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 U
170 }
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/// @endcond
179/// C-order (row-major) strides for a shape.
180inline 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;
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).
191inline 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;
201/// Convert signed dims/indices to sizes, rejecting negatives (a negative cast to
202/// size_t becomes huge -> under-allocation / OOB). Validate at the boundary.
203inline 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;
211/// Advance a C-order multi-index odometer; false when it wraps past the end.
212inline 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;
220/// Peel `std::vector<>` layers off a (possibly deeply nested) list type to reach the
221/// leaf scalar — `nested_scalar_t<std::vector<std::vector<double>>>` is `double`.
222template <typename V> struct nested_scalar { using type = V; };
223template <typename U> struct nested_scalar<std::vector<U>> {
224 using type = typename nested_scalar<U>::type;
225};
226template <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).
229template <typename V> inline constexpr bool is_std_vector_v = false;
230template <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).
233template <Element T>
234void 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));
237/// Walk a nested list: record each axis length the first time it is seen, reject a
238/// ragged list (a row whose length differs from its siblings — numpy does too), and
239/// flatten the leaves in C-order. The leaf scalar must be a @ref Field.
240template <typename U>
241 requires Field<nested_scalar_t<U>>
242void 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);
250} // namespace detail
252/// The backing store of an ndarray: a flat, contiguous, shared element buffer. It uses
253/// @ref detail::default_init_allocator so a freshly-sized result buffer that an op is
254/// 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.
256template <Element T>
257using 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 without
264 * copying elements. Index math goes through @ref at, which bounds-checks. The element
265 * 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 types
267 * (`std::complex<double>`) make complex matrices/vectors — and the complex eigenvalues
268 * a real matrix can have — first-class.
269 */
270template <Element T>
271class basic_ndarray {
272public:
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 this
278 * 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.ToStringScalar
282 * @systest StdlibE2E.Ndarray
283 */
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 and
289 * 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>>`) of
294 * `product(shape)` elements; throws if the shape overflows size_t.
295 * @test CheatahNDArray.ShapeFactoriesAndReductions
296 * @systest StdlibE2E.Ndarray
297 */
298 explicit basic_ndarray(std::vector<std::size_t> shape, T fill = T{}) // contiguous
299 : data_(std::make_shared<buffer_t<T>>()), shape_(std::move(shape)) {
300 // resize (default-init: no zero pass) then std::fill — the fill goes through
301 // 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 a
303 // 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 left
310 * UNINITIALIZED — for internal ops (binary ops, ufuncs, reshape, array) that
311 * 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.BroadcastingAdd
317 * @systest StdlibE2E.Ndarray
318 */
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-fill
323 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 new
330 * array shares ownership of @p data; callers (e.g. @ref broadcast_to) are
331 * 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.BroadcastTo
339 * @systest StdlibE2E.Ndarray
340 */
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.ShapeFactoriesAndReductions
352 * @systest StdlibE2E.Ndarray
353 */
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.BroadcastTo
361 */
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.BroadcastingAdd
369 */
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 than
375 * 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.ShapeFactoriesAndReductions
380 */
381 std::size_t size() const { return detail::product(shape_); } // 1 for a 0-d scalar
383 /**
384 * MUTABLE element reference by multi-index — the write path behind cheatah
385 * subscript assignment `x[i] = v` / `x[i, j] = v`. Negative indices count
386 * 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.SubscriptReadWrite
392 * @crtest LangFeatures.NdarraySubscript
393 */
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 as
402 * 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.SubscriptReadWrite
408 * @crtest LangFeatures.NdarraySubscript
409 */
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 scoped
417 /// 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.SubscriptReadWrite
422 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.SubscriptReadWrite
433 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 a
452 * 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.SliceAssignCopiesIn
458 * @systest StdlibE2E.Ndarray
459 */
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 via
473 * the array's strides. Computes the flat buffer position as
474 * `offset + sum(index[i] * strides[i])`, so it correctly resolves views (including
475 * 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.RejectsMaliciousShapesAndIndices
483 */
484 T at(const std::vector<std::size_t>& index) const { // element via strides
485 // Bounds-check: a wrong-rank or out-of-range index would otherwise compute an
486 // 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.BroadcastingAdd
503 */
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.BroadcastTo
511 */
512 std::size_t offset() const { return offset_; }
514 /**
515 * Python-style text rendering, e.g. `"[[1, 2], [3, 4]]"` — the `str()` member, the
516 * same full form `to_string` and `operator<<` produce (`io.print` reaches an array
517 * through @ref cheatah_pretty_print and `operator<<`, not this hook). Defers to
518 * 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.PrettyPrintAbbreviatesLarge
523 */
524 std::string str() const;
526 /**
527 * Pretty-print hook used by `io.print`: renders the array in nested-bracket form but
528 * ABBREVIATES a large array with `...` (numpy-style edge items), so printing a big array
529 * 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.PrettyPrintAbbreviatesLarge
537 * @crtest NdarrayCompileRun.PrintAbbreviatesLargeArray
538 */
539 void cheatah_pretty_print(std::ostream& os, long long indent) const;
541private:
542 std::shared_ptr<buffer_t<T>> data_;
543 std::vector<std::size_t> shape_;
544 std::vector<std::ptrdiff_t> strides_; // element strides
545 std::size_t offset_ = 0;
546};
548/// The default ndarray element type is `double` — `NDArray` names that
549/// instantiation (the std::string ↔ std::basic_string<char> pattern), so existing
550/// code and the linalg routines keep working unchanged.
551using 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 missing
557 * leading dims as 1; each output dim is the non-1 input dim, and two unequal dims
558 * 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.BroadcastShapeRules
565 * @systest StdlibE2E.Ndarray
566 */
567std::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 data
572// (array([1,2,3]) -> long long, array([1.0,…]) -> double); every op is
573// constrained by Numeric. Elementwise ops vectorize via the std::execution
574// policies (declarative SIMD); broadcasting/strided views fall back to a
575// 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.BroadcastingAdd
586 * @systest StdlibE2E.Ndarray
587 */
588template <Element T>
589inline bool is_contiguous(const basic_ndarray<T>& a) {
590 // C-order contiguity WITHOUT materializing the reference strides: walk the dims
591 // back-to-front and check each stride equals the running size product. The old
592 // `strides() == contiguous_strides(shape())` heap-allocated a vector on every call
593 // — a fixed cost that dominated small-n reductions (e.g. dot, where it ran twice
594 // 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;
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 the
612 * shapes are not broadcast-compatible.
613 * @test CheatahNDArray.BroadcastTo
614 * @systest StdlibE2E.Ndarray
615 */
616template <Element T>
617basic_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 0
621 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());
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 type
636 * (`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.Arange
642 * @crtest NdarrayCompileRun.Arange
643 * @systest StdlibE2E.Ndarray
644 */
645template <Copyable T>
646 requires (!detail::is_std_vector_v<T>) // a vector-of-vectors is a NESTED list (overload below)
647basic_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;
652/**
653 * 1-D array that MOVES its elements out of @p values into a fresh buffer (no element copy) — the
654 * no-copy build path, and the ONLY `array` overload a move-only element type has. A temporary
655 * `array(std::vector<T>{…})` binds here automatically; a named lvalue you want to keep uses the
656 * 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.ArrayMoveIn
662 */
663template <Element T>
664 requires (!detail::is_std_vector_v<T>) // a vector-of-vectors is a NESTED list (overload below)
665basic_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 buffer
668 std::move(src.begin(), src.end(), a.buffer()->begin());
669 return a;
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 off
674 * the nesting (outer list = axis 0, …) and the leaf scalar type is deduced; the list
675 * must be **rectangular** (every sibling row the same length) or it throws, exactly as
676 * numpy rejects a ragged array. Selected only when the argument is itself a list of
677 * 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 leaf
679 * 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.NestedArrayConstruction
685 * @crtest NdarrayCompileRun.NestedArray
686 */
687template <typename V>
688 requires detail::is_std_vector_v<V> && Element<detail::nested_scalar_t<V>>
689basic_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;
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.ShapeFactoriesAndReductions
706 */
707template <Copyable T>
708 requires (!detail::is_std_vector_v<T>) // nested braces route to the nested-list overload
709basic_ndarray<T> array(std::initializer_list<T> values) {
710 return array(std::vector<T>(values));
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.ElementwiseAndScalarBroadcast
719 * @crtest NdarrayCompileRun.Scalar
720 * @systest StdlibE2E.Ndarray
721 */
722template <Copyable T>
723basic_ndarray<T> scalar(T value) {
724 basic_ndarray<T> a; // 0-d
725 a.buffer()->assign(1, value);
726 return a;
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.ShapeFactoriesAndReductions
735 * @crtest NdarrayCompileRun.Zeros
736 * @systest StdlibE2E.Ndarray
737 */
738inline NDArray zeros(const std::vector<long long>& shape) {
739 return NDArray(detail::to_size(shape), 0.0);
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.ShapeFactoriesAndReductions
748 * @crtest NdarrayCompileRun.Ones
749 * @systest StdlibE2E.Ndarray
750 */
751inline NDArray ones(const std::vector<long long>& shape) {
752 return NDArray(detail::to_size(shape), 1.0);
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.RejectsMaliciousShapesAndIndices
762 * @crtest NdarrayCompileRun.Full
763 * @systest StdlibE2E.Ndarray
764 */
765template <Copyable T>
766basic_ndarray<T> full(const std::vector<long long>& shape, T value) {
767 return basic_ndarray<T>(detail::to_size(shape), value);
769/**
770 * A fresh array with the SAME shape and element type as @p a, filled with @p value
771 * (≈ `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.LikeFactories
778 */
779template <Copyable T>
780basic_ndarray<T> full_like(const basic_ndarray<T>& a, T value) {
781 return basic_ndarray<T>(a.shape(), value);
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.LikeFactories
791 */
792template <Copyable T>
793basic_ndarray<T> zeros_like(const basic_ndarray<T>& a) {
794 return basic_ndarray<T>(a.shape(), T{});
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.LikeFactories
803 */
804template <Copyable T>
805basic_ndarray<T> ones_like(const basic_ndarray<T>& a) {
806 return basic_ndarray<T>(a.shape(), T{1});
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.Arange
818 * @crtest NdarrayCompileRun.Arange
819 * @systest StdlibE2E.Ndarray
820 */
821template <Numeric T>
822basic_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 every
828 // pass; the multiply form keeps each element one rounding away from exact. The
829 // 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);
843/**
844 * Reshape @p a to @p shape (same element count); reads in C-order so views/broadcasts
845 * 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.ReshapeSizeMismatchThrows
852 * @crtest NdarrayCompileRun.Reshape
853 * @systest StdlibE2E.Ndarray
854 */
855template <Copyable T>
856basic_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 below
862 auto& buf = *out.buffer();
863 // Contiguous source (the common case — e.g. reshaping a freshly built array): copy
864 // the flat block in one shot instead of walking a per-element bounds-checked
865 // 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 element
874 buf[flat++] = a.at(idx);
875 if (a.ndim() == 0 || !detail::next_index(idx, a.shape())) break;
876 }
877 return out;
879/**
880 * Convert @p a to a new array with element type @p U — numpy's `a.astype(dtype)`. Every element
881 * is `static_cast` into @p U, so this is the way to build a NARROW-element array (a smaller memory
882 * 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 — a
884 * copy, never an alias), same shape out as in. Widening is exact; narrowing truncates/wraps at the
885 * target width (as in C / a numpy fixed dtype). Constrained to conversions that actually exist
886 * (`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.AstypeNarrowsAndWidens
893 * @crtest NdarrayCompileRun.Astype
894 * @systest StdlibE2E.Ndarray
895 */
896template <Field U, Field T>
897 requires std::convertible_to<T, U>
898basic_ndarray<U> astype(const basic_ndarray<T>& a) {
899 basic_ndarray<U> out = basic_ndarray<U>::uninitialized(a.shape()); // every element is written below
900 auto& buf = *out.buffer();
901 if (is_contiguous(a)) { // contiguous source: one straight cast pass, no odometer
902 const T* src = a.buffer()->data() + a.offset();
903 // NOLINT below: an i8→wider astype must sign-extend (numpy dtype semantics) — these
904 // 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 element
911 buf[flat++] = static_cast<U>(a.at(idx)); // NOLINT(bugprone-signed-char-misuse,cert-str34-c): same sign-extension intent as above
912 if (a.ndim() == 0 || !detail::next_index(idx, a.shape())) break;
913 }
914 return out;
917// ---- element-wise ops (broadcasting, vectorized) ----
918/// @cond INTERNAL
919/// 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 a
922 * fresh contiguous result. Fast path: when both operands are contiguous, a flat
923 * `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.BroadcastingAdd
931 * @systest StdlibE2E.Ndarray
932 */
933template <Field T, typename Op>
934basic_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 below
937 auto& obuf = *out.buffer();
938 // Scalar fast paths: `array ⊕ scalar` (or the reverse) is by far the most common
939 // broadcast, and the general strided walk below does a bounds-checked at() per
940 // element (no SIMD). When the other operand is a single value over a contiguous
941 // full-shape array, it's a flat loop we hand to the unseq transform so it vectorizes
942 // 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 element
970 obuf[flat++] = op(av.at(idx), bv.at(idx));
971 if (rshape.empty() || !detail::next_index(idx, rshape)) break;
972 }
973 return out;
976/**
977 * Elementwise `out = op(a, b)` (broadcasting) into the CALLER'S buffer @p out — the user-provided-output
978 * form of binary_op, NO allocation: a hot loop hands the same scratch array every call. @p out must
979 * already hold the broadcast result shape and be contiguous; it MAY alias a full-shape operand (the write
980 * 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.BinaryOpIntoReusesBuffer
988 */
989template <Field T, typename Op>
990void 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 ⊕ scalar
999 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;
1005 if (a.size() == 1 && b.shape() == rshape && is_contiguous(b)) { // scalar ⊕ array
1006 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;
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 transform
1015 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;
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 element
1023 odst[static_cast<std::ptrdiff_t>(flat++)] = op(av.at(idx), bv.at(idx));
1024 if (rshape.empty() || !detail::next_index(idx, rshape)) break;
1027/// @endcond
1029// Shared elementwise combiners: ONE functor type per op, used by BOTH the allocating
1030// forms (add/sub/mul/divide) and the in-place compound operators (+=/-=/*=//=). Using a
1031// single type means `binary_op` is instantiated once per op rather than once per call
1032// site, so the in-place fallback reuses the same (already-tested) instantiation instead
1033// of a duplicate whose scalar/contiguous fast paths are unreachable through it.
1034namespace detail {
1035struct add_op { template <typename T> T operator()(T x, T y) const { return x + y; } };
1036struct sub_op { template <typename T> T operator()(T x, T y) const { return x - y; } };
1037struct mul_op { template <typename T> T operator()(T x, T y) const { return x * y; } };
1038struct 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. compute
1040// `dst = src OP dst` so that `a - std::move(b)` / `a / std::move(b)` can write through b's buffer.
1041struct rsub_op { template <typename T> T operator()(T x, T y) const { return y - x; } };
1042struct rdiv_op { template <typename T> T operator()(T x, T y) const { return y / x; } };
1043} // namespace detail
1045/**
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.BroadcastingAdd
1052 * @crtest NdarrayCompileRun.Add
1053 * @systest StdlibE2E.Ndarray
1054 */
1055template <Field T>
1056basic_ndarray<T> add(const basic_ndarray<T>& a, const basic_ndarray<T>& b) {
1057 return binary_op(a, b, detail::add_op{});
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.ElementwiseAndScalarBroadcast
1066 * @crtest NdarrayCompileRun.Sub
1067 * @systest StdlibE2E.Ndarray
1068 */
1069template <Field T>
1070basic_ndarray<T> sub(const basic_ndarray<T>& a, const basic_ndarray<T>& b) {
1071 return binary_op(a, b, detail::sub_op{});
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.ElementwiseAndScalarBroadcast
1080 * @crtest NdarrayCompileRun.Mul
1081 * @systest StdlibE2E.Ndarray
1082 */
1083template <Field T>
1084basic_ndarray<T> mul(const basic_ndarray<T>& a, const basic_ndarray<T>& b) {
1085 return binary_op(a, b, detail::mul_op{});
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.ElementwiseAndScalarBroadcast
1094 * @crtest NdarrayCompileRun.Divide
1095 * @systest StdlibE2E.Ndarray
1096 */
1097template <Field T>
1098basic_ndarray<T> divide(const basic_ndarray<T>& a, const basic_ndarray<T>& b) {
1099 return binary_op(a, b, detail::div_op{});
1102/// @cond INTERNAL
1103/// 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 a
1107 * hot loop hands the same scratch every call. @p out must be contiguous with the broadcast shape; it
1108 * 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.BinaryOpIntoReusesBuffer
1113 */
1114template <Field T>
1115void 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{});
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 */
1123template <Field T>
1124void 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{});
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 */
1132template <Field T>
1133void 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{});
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 */
1141template <Field T>
1142void 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{});
1145/// @endcond
1147// ---- Infix operators & in-place compound assignment ------------------------
1148// cheatah lowers `a + b` / `a * 2.0` / `a += b` on ndarrays straight to these
1149// C++ operators. Infix forms are the elementwise free functions (broadcasting
1150// 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 common
1152// (contiguous) layout — no allocation, so a hot loop can reuse one array for
1153// an entire run — falling back to the allocating elementwise path only for
1154// non-contiguous views or true broadcasts.
1156/// @cond INTERNAL
1157/// 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 as
1161 * a flat vectorizable transform with NO allocation; anything else falls back
1162 * 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.CompoundAssignInPlace
1169 * @crtest LangFeatures.NdarrayOperators
1170 * @systest StdlibE2E.Ndarray
1171 */
1172template <typename T, typename Op>
1173void 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;
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;
1191 a = binary_op(a, b, op); // broadcast / non-contiguous fallback
1193/// @endcond
1195/// @cond INTERNAL
1196/// 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 clobber
1199// `a`, which the caller still holds. But when the LEFT operand is an RVALUE — a temporary the caller
1200// has already given up: the `a + b` inside a chain `a + b + c`, or an explicit `std::move(a)` — these
1201// compute IN PLACE into that buffer and move it out: NO allocation. Selected by value category, so a
1202// buffer is only ever reused when it is safe to (no flag, no surprise mutation). Reuses the in-place
1203// 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 */
1211template <Field T>
1212basic_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);
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 */
1223template <Field T>
1224basic_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);
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 */
1235template <Field T>
1236basic_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);
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 */
1247template <Field T>
1248basic_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);
1253// Right-operand reuse: when only the RIGHT operand is the expiring temporary, compute through ITS
1254// buffer instead. `+`/`*` are commutative so `op(b, a)` is the same value; `-`/`/` use the reversed
1255// combiners (`b = a OP b`). This makes `a + std::move(b)` reuse a buffer exactly like
1256// `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 */
1264template <Field T>
1265basic_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);
1269/**
1270 * Element-wise `a - b` reusing the expiring right operand @p b in place via the reversed combiner
1271 * `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 */
1277template <Field T>
1278basic_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);
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 */
1289template <Field T>
1290basic_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);
1294/**
1295 * Element-wise `a / b` reusing the expiring right operand @p b in place via the reversed combiner
1296 * `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 */
1302template <Field T>
1303basic_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);
1308// Both operands expiring: prefer reusing the LEFT (matches the chain `a + b + c`, where the left is
1309// 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 */
1317template <Field T>
1318basic_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 */
1326template <Field T>
1327basic_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 */
1335template <Field T>
1336basic_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 */
1344template <Field T>
1345basic_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/// @endcond
1348/// 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.RvalueOperandReusesBuffer
1354 */
1355template <typename T>
1356basic_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 */
1362template <typename T>
1363basic_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 */
1369template <typename T>
1370basic_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.DivideInfixLvalueForm
1376 */
1377template <typename T>
1378basic_ndarray<T> operator/(const basic_ndarray<T>& a, const basic_ndarray<T>& b) { return divide(a, b); }
1380/// @cond INTERNAL
1381/// 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 PLACE
1383/// (no alloc). `std::move(a) + b` reuses `a`, `a + std::move(b)` reuses `b` — symmetric. A chain
1384/// `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. */
1386template <typename T>
1387basic_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. */
1389template <typename T>
1390basic_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. */
1392template <typename T>
1393basic_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. */
1395template <typename T>
1396basic_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. */
1398template <typename T>
1399basic_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. */
1401template <typename T>
1402basic_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. */
1404template <typename T>
1405basic_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. */
1407template <typename T>
1408basic_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. */
1410template <typename T>
1411basic_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. */
1413template <typename T>
1414basic_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. */
1416template <typename T>
1417basic_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. */
1419template <typename T>
1420basic_ndarray<T> operator/(basic_ndarray<T>&& a, basic_ndarray<T>&& b) { return divide(std::move(a), std::move(b)); }
1421/// @endcond
1423/// 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 go
1425/// straight to the ALLOCATING binary_op (not the reuse-enabled add/sub/...): the `scalar(s)` temporary
1426/// is 0-d, so letting it bind a buffer-reuse overload would compute the result into the scalar and
1427/// collapse it to 0-d. The array operand here is a const lvalue (the caller keeps it), so the result
1428/// 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 */
1437template <typename T, typename S>
1438 requires std::is_arithmetic_v<S>
1439basic_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 */
1448template <typename T, typename S>
1449 requires std::is_arithmetic_v<S>
1450basic_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 */
1459template <typename T, typename S>
1460 requires std::is_arithmetic_v<S>
1461basic_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.ScalarTimesSizeOneArrayKeepsShape
1470 */
1471template <typename T, typename S>
1472 requires std::is_arithmetic_v<S>
1473basic_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.CompoundAssignInPlace
1482 * @crtest LangFeatures.NdarrayOperators
1483 */
1484template <typename T, typename S>
1485 requires std::is_arithmetic_v<S>
1486basic_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.ScalarTimesSizeOneArrayKeepsShape
1495 * @crtest LangFeatures.NdarrayOperators
1496 */
1497template <typename T, typename S>
1498 requires std::is_arithmetic_v<S>
1499basic_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 */
1508template <typename T, typename S>
1509 requires std::is_arithmetic_v<S>
1510basic_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 */
1519template <typename T, typename S>
1520 requires std::is_arithmetic_v<S>
1521basic_ndarray<T> operator/(S s, const basic_ndarray<T>& a) { return binary_op(scalar(static_cast<T>(s)), a, detail::div_op{}); }
1523/// @cond INTERNAL
1524/// 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 the
1526/// commutative `s + a` / `s * a` get a scalar-LEFT reuse form; `s - a` / `s / a` keep the allocating
1527/// 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. */
1529template <typename T, typename S>
1530 requires std::is_arithmetic_v<S>
1531basic_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. */
1533template <typename T, typename S>
1534 requires std::is_arithmetic_v<S>
1535basic_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. */
1537template <typename T, typename S>
1538 requires std::is_arithmetic_v<S>
1539basic_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. */
1541template <typename T, typename S>
1542 requires std::is_arithmetic_v<S>
1543basic_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. */
1545template <typename T, typename S>
1546 requires std::is_arithmetic_v<S>
1547basic_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. */
1549template <typename T, typename S>
1550 requires std::is_arithmetic_v<S>
1551basic_ndarray<T> operator/(basic_ndarray<T>&& a, S s) { return divide(std::move(a), scalar(static_cast<T>(s))); }
1552/// @endcond
1554/// 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.CompoundAssignInPlace
1564 * @test CheatahNDArray.CompoundAssignNonContiguousFallback
1565 * @crtest LangFeatures.NdarrayOperators
1566 */
1567template <typename T>
1568basic_ndarray<T>& operator+=(basic_ndarray<T>& a, const basic_ndarray<T>& b) {
1569 compound_apply(a, b, detail::add_op{}); return a;
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.CompoundAssignNonContiguousFallback
1579 */
1580template <typename T>
1581basic_ndarray<T>& operator-=(basic_ndarray<T>& a, const basic_ndarray<T>& b) {
1582 compound_apply(a, b, detail::sub_op{}); return a;
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.CompoundAssignNonContiguousFallback
1592 */
1593template <typename T>
1594basic_ndarray<T>& operator*=(basic_ndarray<T>& a, const basic_ndarray<T>& b) {
1595 compound_apply(a, b, detail::mul_op{}); return a;
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.CompoundAssignNonContiguousFallback
1605 */
1606template <typename T>
1607basic_ndarray<T>& operator/=(basic_ndarray<T>& a, const basic_ndarray<T>& b) {
1608 compound_apply(a, b, detail::div_op{}); return a;
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 */
1618template <typename T, typename S>
1619 requires std::is_arithmetic_v<S>
1620basic_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.CompoundAssignInPlace
1629 */
1630template <typename T, typename S>
1631 requires std::is_arithmetic_v<S>
1632basic_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.CompoundAssignInPlace
1641 * @crtest LangFeatures.ParamsPassByReference
1642 */
1643template <typename T, typename S>
1644 requires std::is_arithmetic_v<S>
1645basic_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.CompoundAssignInPlace
1654 * @crtest LangFeatures.NdarrayOperators
1655 */
1656template <typename T, typename S>
1657 requires std::is_arithmetic_v<S>
1658basic_ndarray<T>& operator/=(basic_ndarray<T>& a, S s) { return a /= scalar(static_cast<T>(s)); }
1660// ---- complex support ----
1661namespace detail {
1662/// Map @p a element-wise through @p f into a fresh contiguous array of element type
1663/// `U` (which may differ from `T` — e.g. complex→real for @ref real). Contiguous
1664/// fast path via `std::transform(unseq)`; otherwise a C-order walk.
1665template <typename U, Field T, typename F>
1666basic_ndarray<U> map_array(const basic_ndarray<T>& a, F f) {
1667 basic_ndarray<U> out = basic_ndarray<U>::uninitialized(a.shape()); // fully written below
1668 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;
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 element
1678 obuf[flat++] = f(a.at(idx));
1679 if (a.ndim() == 0 || !next_index(idx, a.shape())) break;
1681 return out;
1684// Out-of-line, separately-compiled (-ffast-math) double-precision SIMD kernels for the
1685// element-wise ufuncs — see ufunc_simd.cpp. They vectorize the transcendentals through
1686// libmvec, which the default flags cannot; isolating -ffast-math to that file keeps the
1687// rest of cheatah's arithmetic strict.
1688void simd_sqrt_f64(const double*, double*, std::size_t);
1689void simd_cbrt_f64(const double*, double*, std::size_t);
1690void simd_exp_f64(const double*, double*, std::size_t);
1691void simd_log_f64(const double*, double*, std::size_t);
1692void simd_sin_f64(const double*, double*, std::size_t);
1693void simd_cos_f64(const double*, double*, std::size_t);
1694void simd_tan_f64(const double*, double*, std::size_t);
1696/// Map a ufunc over @p a: a *contiguous double* array goes through the precompiled SIMD
1697/// @p kernel; everything else (float, or a strided/broadcast view) uses the generic
1698/// scalar @p fallback. Same result either way — the kernel just vectorizes the hot case.
1699template <FloatingPoint T, class Kernel, class Fallback>
1700basic_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 it
1704 kernel(a.buffer()->data() + a.offset(), out.buffer()->data(), a.size());
1705 return out;
1708 return map_array<T>(a, fallback);
1710} // namespace detail
1712/**
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/vector
1715 * (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.ComplexConstructAndParts
1722 * @crtest NdarrayCompileRun.Complex
1723 * @systest StdlibE2E.NdarrayComplex
1724 */
1725template <FloatingPoint T>
1726basic_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 element
1736 obuf[flat++] = C(rv.at(idx), iv.at(idx));
1737 if (rshape.empty() || !detail::next_index(idx, rshape)) break;
1739 return out;
1741/**
1742 * Element-wise complex conjugate (`a − b·j` for each `a + b·j`); on a real array it
1743 * is the identity (a copy). Type-preserving. Used to form Hermitian adjoints and
1744 * 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.ComplexConstructAndParts
1750 * @crtest NdarrayCompileRun.Conj
1751 * @systest StdlibE2E.NdarrayComplex
1752 */
1753template <Field T>
1754basic_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;
1761 });
1763/**
1764 * The real parts as a real array (the identity on a real array). For `a + b·j` it
1765 * 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.ComplexConstructAndParts
1771 * @crtest NdarrayCompileRun.Real
1772 * @systest StdlibE2E.NdarrayComplex
1773 */
1774template <Field T>
1775basic_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;
1783 });
1785/**
1786 * The imaginary parts as a real array (all zeros for a real array). For `a + b·j` it
1787 * 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.ComplexConstructAndParts
1793 * @crtest NdarrayCompileRun.Imag
1794 * @systest StdlibE2E.NdarrayComplex
1795 */
1796template <Field T>
1797basic_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};
1805 });
1808// ---- element-wise math (numpy-style ufuncs) ----
1809// These are the array counterparts of the scalar `math` module — mirroring Python's
1810// split: `math.sqrt(x)` for a scalar, `ndarray.sqrt(a)` (≈ `numpy.sqrt`) for a whole
1811// array. A contiguous `double` array routes through a precompiled SIMD kernel
1812// (ufunc_simd.cpp) that vectorizes via glibc's libmvec — so `exp`/`sin`/… run at vector
1813// speed and beat NumPy's ufuncs; other element types / strided views fall back to a
1814// 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.ElementwiseMath
1821 * @crtest NdarrayCompileRun.Sqrt
1823 */
1824template <FloatingPoint T>
1825basic_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); });
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.ElementwiseMath
1835 */
1836template <FloatingPoint T>
1837basic_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); });
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.ElementwiseMath
1846 * @crtest NdarrayCompileRun.Exp
1848 */
1849template <FloatingPoint T>
1850basic_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); });
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.ElementwiseMath
1859 */
1860template <FloatingPoint T>
1861basic_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); });
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.ElementwiseMath
1870 * @crtest NdarrayCompileRun.Sin
1871 * @systest StdlibE2E.NdarrayMath
1872 */
1873template <FloatingPoint T>
1874basic_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); });
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.ElementwiseMath
1883 */
1884template <FloatingPoint T>
1885basic_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); });
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.ElementwiseMath
1894 */
1895template <FloatingPoint T>
1896basic_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); });
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.ElementwiseMath
1905 * @systest StdlibE2E.NdarrayMath
1906 */
1907template <FloatingPoint T>
1908basic_ndarray<T> abs(const basic_ndarray<T>& a) {
1909 return detail::map_array<T>(a, [](T x) { return std::fabs(x); });
1912// ---- reductions / access / display ----
1913namespace detail {
1914/// The shared multi-accumulator reduction: sums `get(0)..get(n-1)` with EIGHT independent
1915/// accumulators, tree-combined, plus a scalar tail. The independent lanes break the FP-add
1916/// dependency chain so -O3 -march=native emits SIMD+FMA and reaches memory bandwidth instead of
1917/// add latency (a single running sum — or a plain `std::reduce`, which libstdc++ left-folds for FP
1918/// without -ffast-math — serializes: the dot/norm mistake). `get(i)` returns the i-th TERM — an
1919/// element for `sum`, a (possibly conjugated) product for `dot`, a strided read for `trace`. One
1920/// primitive replaces the copies formerly hand-rolled in ndarray/linalg. `constexpr`, so a
1921/// fixed-extent caller gets a compile-time reduction too.
1922template <class T, class Get>
1923constexpr 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);
1930 T s = ((s0 + s1) + (s2 + s3)) + ((s4 + s5) + (s6 + s7));
1931 for (; i < n; ++i) s += get(i);
1932 return s;
1934} // namespace detail
1935/**
1936 * Sum of all elements — a full reduction across every axis (a contiguous array goes
1937 * through the shared multi-accumulator SIMD reduction @ref detail::reduce_lanes, else
1938 * 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.ShapeFactoriesAndReductions
1944 * @crtest NdarrayCompileRun.Sum
1945 * @systest StdlibE2E.Ndarray
1946 */
1947template <Field T>
1948T 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]; });
1954 T s{};
1955 std::vector<std::size_t> idx(a.ndim(), 0);
1956 for (;;) { // a 0-d array still contributes its one element
1957 s += a.at(idx);
1958 if (a.ndim() == 0 || !detail::next_index(idx, a.shape())) break;
1960 return s;
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.ShapeFactoriesAndReductions
1969 * @crtest NdarrayCompileRun.Mean
1970 * @systest StdlibE2E.Ndarray
1971 */
1972template <Numeric T>
1973double 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);
1977/**
1978 * Read one element by signed multi-index (the cheatah-facing wrapper over @ref
1979 * 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.ShapeFactoriesAndReductions
1986 * @crtest NdarrayCompileRun.Get
1987 * @systest StdlibE2E.Ndarray
1988 */
1989template <Copyable T>
1990T get(const basic_ndarray<T>& a, const std::vector<long long>& index) {
1991 return a.at(detail::to_size(index));
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.ShapeFactoriesAndReductions
2000 * @crtest NdarrayCompileRun.ShapeOf
2001 * @systest StdlibE2E.Ndarray
2002 */
2003template <Element T>
2004std::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;
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.ShapeFactoriesAndReductions
2016 * @crtest NdarrayCompileRun.SizeOf
2017 * @crtest NdarrayCompileRun.RprintShowsArrayFull
2018 */
2019template <Element T>
2020long long size_of(const basic_ndarray<T>& a) {
2021 return static_cast<long long>(a.size());
2024namespace detail {
2025/// Format one element. Real types go through `operator<<`; a complex element is
2026/// rendered Python-style as `a+bj` / `a-bj` (not the `std::complex` default
2027/// `(a,b)`), so a complex spectrum prints the way a cheatah user expects.
2028template <typename T>
2029std::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";
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 character
2044 } else {
2045 os << v;
2047 return os.str();
2050/// Recursively format @p a into nested brackets (each element via `format_scalar`).
2051template <Element T>
2052void 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;
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);
2064 out += "]";
2067/// Like format_rec but ABBREVIATES a large array: an axis longer than `2*edge` shows its
2068/// first and last `edge` items with `...` between, recursively. Summarization is enabled by
2069/// @p summarize (the caller turns it on only past a total-size threshold), so small arrays
2070/// print in full.
2071template <Element T>
2072void 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;
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;
2089 continue; // skip the abbreviated middle
2091 if (!first) out += ", ";
2092 first = false;
2093 idx[dim] = i;
2094 format_rec_trunc(a, idx, dim + 1, out, edge, summarize);
2096 out += "]";
2099/// to_string, but ABBREVIATED with `...` when the array is large (total size beyond a
2100/// numpy-style threshold) — the readable default for `io.print`. `io.rprint`/`to_string`
2101/// keep the full form.
2102template <Element T>
2103std::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 axis
2106 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;
2112} // namespace detail
2114/**
2115 * Render as a nested-bracket string, e.g. `"[[1, 2], [3, 4]]"` (a 0-d scalar renders
2116 * 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.BroadcastingAdd
2122 * @crtest NdarrayCompileRun.ToString
2123 * @systest StdlibE2E.Ndarray
2124 */
2125template <Element T>
2126std::string to_string(const basic_ndarray<T>& a) {
2127 if (a.ndim() == 0) {
2128 return detail::format_scalar(a.at({}));
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;
2136template <Element T>
2137inline std::string basic_ndarray<T>::str() const {
2138 return to_string(*this);
2141/**
2142 * Stream an array to a `std::ostream` (the FULL nested-bracket form) — so an NDArray is
2143 * directly Streamable, like a primitive or a cheatah struct, without going through
2144 * `to_string`/`str()`. `io.rprint`, `str()`, and a struct that holds an array all stream it
2145 * 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.StreamableOperator
2152 * @systest StdlibE2E.Ndarray
2153 */
2154template <Element T>
2155std::ostream& operator<<(std::ostream& os, const basic_ndarray<T>& a) {
2156 return os << to_string(a);
2159template <Element T>
2160inline void basic_ndarray<T>::cheatah_pretty_print(std::ostream& os, long long /*unused*/) const {
2161 os << detail::to_string_pretty(*this);
2164} // namespace cheatah::ndarray
2166// 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.
2169namespace cheatah::builtins {
2171/// @cond INTERNAL
2172/// This header deliberately includes no project header — it only REOPENS the builtins namespace to
2173/// add ndarray overloads, exactly as `index` below does. So the two names the slice-assignment
2174/// vocabulary needs are spelled locally: `empty_seq` is forward-declared (a reference to an
2175/// incomplete type is all an overload declaration needs, and builtins.hpp completes it wherever
2176/// both headers are in play), and the "to the end" sentinel is the same value builtins::slice_end
2177/// holds. A static_assert in builtins.hpp keeps the two from drifting.
2178struct empty_seq;
2179inline 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).
2181inline long long nd_norm_index(long long i, long long n) { return i < 0 ? i + n : i; }
2182/// @endcond
2185/**
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.SubscriptReadWrite
2193 * @crtest LangFeatures.NdarraySubscript
2194 * @systest StdlibE2E.Ndarray
2195 */
2196template <typename T, ::cheatah::ndarray::Subscript First, ::cheatah::ndarray::Subscript... Ix>
2197T 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 conversion
2199 // (ndarray::Subscript); the const overload reads, and the value is copied out.
2200 return a.item_ref(first, rest...);
2204/**
2205 * A view of `a[lo:hi]` along axis 0 — the rows `lo` up to `hi`, sharing @p a's buffer.
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 — negatives
2210 * 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.SliceIsAView
2219 * @crtest LangFeatures.NdarraySliceAssignment
2220 * @systest StdlibE2E.Ndarray
2221 */
2222template <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);
2239/**
2240 * Write @p rhs into the elements `a[lo:hi]` addresses — an array assignment COPIES, it never
2241 * rebinds or resizes.
2243 * @p rhs may be a scalar (filling every addressed element) or an array broadcastable to the
2244 * slice's shape, matching numpy. The destination keeps its shape, so an extent that does not
2245 * broadcast is an error rather than a silent truncation. Strides are honoured, so writing
2246 * 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.SliceAssignCopiesIn
2256 * @crtest LangFeatures.NdarraySliceAssignment
2257 * @systest StdlibE2E.Ndarray
2258 */
2259template <typename T, typename R>
2260void 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 fill
2266 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;
2271 } else { // array source, broadcast to the slice
2272 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;
2281/// @cond INTERNAL
2282/// An ndarray has a fixed shape, so there is nothing for `a[lo:hi] = []` to delete.
2283template <typename T>
2284void 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)");
2291/// @endcond
2293} // namespace cheatah::builtins