Source
stdlib/linalg/backend.hpp
1
// Copyright (c) 2026 BigBrain LLC. MIT-licensed (see LICENSE).2
// Original work; see ACKNOWLEDGMENTS.md for the open-source ideas we build upon.3
#pragma once5
/**6
* @file backend.hpp7
* @brief cheatah `linalg` — the two-layer (element T + container Array) generic fronts.8
*9
* Every routine takes TWO template layers: the element `T` and the container template10
* `Array`, with a `requires` concept enforcing the container. Both operands are spelled11
* `Array<T>`, so a host⊗device or f64⊗f32 mix cannot deduce a single `Array`/`T` and is a12
* compile error — the location/element firewall is FREE, via deduction, with no runtime13
* check and no `SameLocation` clause on the common binary ops.14
*15
* Each op is a pair of same-named overloads:16
* - a 2-arg **allocating front** `Array<T> op(const Array<T>&, const Array<T>&)` that17
* allocates the result with `Array<T>::uninitialized(...)` and calls the out-param form;18
* - a 3-arg **out-parameter kernel** `void op(Array<T>& out, …)`, split by concept: the19
* HOST overload (declared here, defined in routines.cpp) runs the raw-pointer SIMD kernel;20
* a device extension supplies a `requires DeviceArray<Array<T>>` overload in ITS namespace,21
* found by ADL. Mutually exclusive concepts → no ambiguity, and cheatah never names the22
* extension. (No CPO objects, no tag_invoke — plain concept-constrained overloads.)23
*/24
#include <stdexcept>25
#include <vector>27
#include "concepts.hpp"29
namespace cheatah::linalg {31
/// @cond INTERNAL32
/// Shared BATCHED shape validation. Both the allocating front and the out-parameter kernel run33
/// it, because either can be called directly: the kernel is declared in this header, documented,34
/// and explicitly instantiated, so "the front already checked" is an assumption it cannot make.35
/// It once did, and a 3-D @p a with a 2-D @p b read b.shape()[2] off the end of a 2-element36
/// shape vector before any check ran.37
template <ndarray::Field T, template <typename> class Array>38
inline void check_matmul_batched(const Array<T>& a, const Array<T>& b) {39
if (a.ndim() != 3 || b.ndim() != 3)40
throw std::runtime_error("linalg: batched matmul expects two 3-D operands");41
if (a.shape()[0] != b.shape()[0])42
throw std::runtime_error("linalg: batched matmul batch-count mismatch");43
if (a.shape()[2] != b.shape()[1])44
throw std::runtime_error("linalg: matmul inner dimension mismatch");45
}46
/// @endcond48
/// @cond INTERNAL49
/// the allocation-free out-parameter kernel (HOST overload; a device50
/// extension adds its own `requires DeviceArray<Array<T>>` overload). Declared here so the51
/// allocating front below can call it; defined + explicitly instantiated in routines.cpp.52
/**53
* Matmul into the CALLER'S buffer @p out (out FIRST) — no result allocation (a hot loop hands54
* the same scratch every call).55
* @tparam T the element type.56
* @tparam Array the (host) container template.57
* @param out contiguous [a.rows, b.cols] destination, overwritten; must NOT alias @p a or @p b.58
* @param a,b the operands.59
* @complexity O(n³) (× B for a batch).60
* @alloc none for contiguous operands (product written straight into @p out); a61
* non-contiguous operand is packed once into scratch.62
* @test LinalgRoutines.MatmulIntoReusesBuffer63
* @test LinalgRoutines.ComplexMatmulIntoReusesBuffer64
*/65
template <ndarray::Field T, template <typename> class Array>66
requires HostArray<Array<T>>67
void matmul(Array<T>& out, const Array<T>& a, const Array<T>& b);68
/// @endcond70
/**71
* Matrix multiply — the allocating front. Both operands are `Array<T>` (so host⊗device / element72
* mixes fail to deduce and are compile errors); requires both to be 2-D with matching inner73
* dimensions — or both 3-D for the BATCHED product `[B,M,K] @ [B,K,N] → [B,M,N]` (equal batch74
* counts, strict: no broadcast batching). Allocates the result via `Array<T>::uninitialized` and75
* fills it through the out-parameter kernel — the host SIMD path, or a device shader when76
* `Array` is a device container.77
* @tparam T the element type (`double` / `std::complex<double>`).78
* @tparam Array the container template.79
* @param a m×k matrix, or a B×m×k batch of matrices.80
* @param b k×p matrix, or a B×k×p batch.81
* @return m×p product (or the B×m×p batch), an `Array<T>` of the same container and element.82
* @complexity O(n³) (× B for a batch).83
* @alloc allocates only the result; operands read in place (a strided host view packs once).84
* @concurrency deliberately single-threaded (the fastest-per-core contract); parallelize85
* across independent products in the caller.86
* @test LinalgRoutines.ProductsAndTrace87
* @test LinalgRoutines.BatchedMatmul88
* @crtest LinalgCompileRun.Matmul89
* @systest StdlibE2E.Linalg90
*/91
template <ndarray::Field T, template <typename> class Array>92
requires NumericArray<Array<T>>93
[[nodiscard]] Array<T> matmul(const Array<T>& a, const Array<T>& b) {94
if (a.ndim() == 3 || b.ndim() == 3) {95
check_matmul_batched<T, Array>(a, b);96
Array<T> out = Array<T>::uninitialized({a.shape()[0], a.shape()[1], b.shape()[2]});97
matmul(out, a, b);98
return out;99
}100
if (a.ndim() != 2 || b.ndim() != 2)101
throw std::runtime_error("linalg: matmul expects 2-D matrices");102
if (a.shape()[1] != b.shape()[0])103
throw std::runtime_error("linalg: matmul inner dimension mismatch");104
Array<T> out = Array<T>::uninitialized({a.shape()[0], b.shape()[1]});105
matmul(out, a, b);106
return out;107
}109
/**110
* The flattened length of a vector-shaped operand — 1-D, or 2-D with a size-1 row/column111
* (throws otherwise). Reads only host-resident shape metadata, so it is valid for ANY located112
* container, device arrays included; the shared validation step of every vector front below.113
* @tparam A the (located) container type.114
* @param a the operand whose vector length is wanted.115
* @return the element count of the flattened vector.116
* @complexity O(1).117
* @alloc none.118
* @test LinalgRoutines.ProductsAndTrace119
*/120
template <NumericArray A>121
[[nodiscard]] inline std::size_t vector_len(const A& a) {122
if (a.ndim() == 1) return a.shape()[0];123
if (a.ndim() == 2 && (a.shape()[0] == 1 || a.shape()[1] == 1)) return a.size();124
throw std::runtime_error("linalg: expected a 1-D vector");125
}127
// ---- reductions (dot / vdot / inner / trace): the scalar-out kernel pattern ----128
// Same two-layer seam as matmul, with a SCALAR out-parameter: the front validates and calls the129
// unqualified `op(out, …)`, which resolves to the HOST kernel below (routines.cpp) or a device130
// extension's `requires DeviceArray<Array<T>>` overload via ADL.132
/// @cond INTERNAL133
/// the scalar-out reduction kernels (HOST overloads; a device extension adds its134
/// own `requires DeviceArray<Array<T>>` overloads). Declared here so the allocating fronts below135
/// can call them; defined + explicitly instantiated in routines.cpp.136
/**137
* Bilinear dot product Σ aᵢbᵢ into the caller's scalar @p out (out FIRST) — the reduction analogue138
* of the out-param matmul kernel, shared by real and complex elements.139
* @tparam T the element type; @tparam Array the (host) container template.140
* @param out receives the scalar sum. @param a,b same-length vectors (validated by the front).141
* @test LinalgRoutines.ProductsAndTrace142
*/143
template <ndarray::Field T, template <typename> class Array>144
requires HostArray<Array<T>>145
void dot(T& out, const Array<T>& a, const Array<T>& b);146
/**147
* Hermitian inner product Σ conj(aᵢ)·bᵢ into @p out (bilinear for a real element — the148
* conjugation is an `if constexpr` branch in the host kernel).149
* @tparam T the element type; @tparam Array the (host) container template.150
* @param out receives the scalar sum. @param a,b same-length vectors (validated by the front).151
* @test LinalgRoutines.VdotInnerOuterKron152
*/153
template <ndarray::Field T, template <typename> class Array>154
requires HostArray<Array<T>>155
void vdot(T& out, const Array<T>& a, const Array<T>& b);156
/**157
* Bilinear inner product Σ aᵢbᵢ into @p out (numpy's `inner`; identical to @ref dot for158
* flattened vectors).159
* @tparam T the element type; @tparam Array the (host) container template.160
* @param out receives the scalar sum. @param a,b same-length vectors (validated by the front).161
* @test LinalgRoutines.VdotInnerOuterKron162
*/163
template <ndarray::Field T, template <typename> class Array>164
requires HostArray<Array<T>>165
void inner(T& out, const Array<T>& a, const Array<T>& b);166
/**167
* Trace (diagonal sum) into @p out — strided diagonal read, no copy.168
* @tparam T the element type; @tparam Array the (host) container template.169
* @param out receives the diagonal sum. @param a a 2-D matrix (validated by the front).170
* @test LinalgRoutines.ProductsAndTrace171
*/172
template <ndarray::Field T, template <typename> class Array>173
requires HostArray<Array<T>>174
void trace(T& out, const Array<T>& a);175
/// @endcond177
/**178
* Dot product: 1-D inner product (vectors flattened) — the bilinear Σ aᵢbᵢ. Flattens each179
* operand to a vector (1-D, or 2-D with a size-1 row/column) and throws if either is not180
* vector-shaped or the lengths differ. Both operands are `Array<T>` (the deduction firewall).181
* @tparam T the element type.182
* @tparam Array the container template.183
* @param a,b same-length vectors.184
* @return Σ aᵢbᵢ as the scalar `T`.185
* @complexity O(n).186
* @alloc none for contiguous operands (read in place); a non-contiguous view packs once O(n).187
* @test LinalgRoutines.ProductsAndTrace188
* @test LinalgRoutines.ComplexProducts189
* @crtest LinalgCompileRun.Dot190
* @crtest LinalgCompileRun.ComplexDot191
* @systest StdlibE2E.Linalg192
* @systest StdlibE2E.LinalgComplex193
*/194
template <ndarray::Field T, template <typename> class Array>195
requires NumericArray<Array<T>>196
[[nodiscard]] T dot(const Array<T>& a, const Array<T>& b) {197
if (vector_len(a) != vector_len(b))198
throw std::runtime_error("linalg: dot dimension mismatch");199
T out;200
dot(out, a, b);201
return out;202
}204
/**205
* Vector dot product. For a REAL element this is the bilinear Σ aᵢbᵢ (identical to @ref dot and206
* @ref inner); for a **complex** element it is the conjugate-linear Hermitian inner product207
* ⟨a, b⟩ = Σ conj(aᵢ)·bᵢ (numpy's `vdot`, conjugating the first argument); `vdot(a, a)` is ‖a‖².208
* @tparam T the element type.209
* @tparam Array the container template.210
* @param a,b same-length vectors.211
* @return Σ aᵢbᵢ (real) or Σ conj(aᵢ)·bᵢ (complex), as the scalar `T`.212
* @complexity O(n).213
* @alloc none for contiguous operands; a non-contiguous view packs once O(n).214
* @test LinalgRoutines.VdotInnerOuterKron215
* @test LinalgRoutines.ComplexProducts216
* @crtest LinalgCompileRun.Vdot217
* @crtest LinalgCompileRun.ComplexVdot218
* @systest StdlibE2E.Linalg219
* @systest StdlibE2E.LinalgComplex220
*/221
template <ndarray::Field T, template <typename> class Array>222
requires NumericArray<Array<T>>223
[[nodiscard]] T vdot(const Array<T>& a, const Array<T>& b) {224
if (vector_len(a) != vector_len(b))225
throw std::runtime_error("linalg: dot dimension mismatch");226
T out;227
vdot(out, a, b);228
return out;229
}231
/**232
* Inner product of two vectors — the bilinear Σ aᵢbᵢ (numpy's `inner`; same as @ref dot).233
* @tparam T the element type.234
* @tparam Array the container template.235
* @param a,b same-length vectors.236
* @return Σ aᵢbᵢ as the scalar `T`.237
* @complexity O(n).238
* @alloc none for contiguous operands; a non-contiguous view packs once O(n).239
* @test LinalgRoutines.VdotInnerOuterKron240
* @crtest LinalgCompileRun.Inner241
* @systest StdlibE2E.Linalg242
*/243
template <ndarray::Field T, template <typename> class Array>244
requires NumericArray<Array<T>>245
[[nodiscard]] T inner(const Array<T>& a, const Array<T>& b) {246
if (vector_len(a) != vector_len(b))247
throw std::runtime_error("linalg: dot dimension mismatch");248
T out;249
inner(out, a, b);250
return out;251
}253
/**254
* Trace: the sum of the matrix diagonal, as the scalar `T`. Requires a 2-D matrix (throws255
* otherwise); rectangular matrices sum min(r, c) diagonal entries.256
* @tparam T the element type.257
* @tparam Array the container template.258
* @param a a 2-D matrix.259
* @return Σ aᵢᵢ as the scalar `T`.260
* @complexity O(min(r, c)).261
* @alloc none (strided diagonal read straight from the buffer).262
* @test LinalgRoutines.ProductsAndTrace263
* @crtest LinalgCompileRun.Trace264
* @systest StdlibE2E.Linalg265
*/266
template <ndarray::Field T, template <typename> class Array>267
requires NumericArray<Array<T>>268
[[nodiscard]] T trace(const Array<T>& a) {269
if (a.ndim() != 2) throw std::runtime_error("linalg: expected a 2-D matrix");270
T out;271
trace(out, a);272
return out;273
}275
// ---- products with array results (outer / conj_transpose / kron): the matmul pattern ----277
/// @cond INTERNAL278
/// the allocation-free out-parameter kernels (HOST overloads; a device extension279
/// adds its own `requires DeviceArray<Array<T>>` overloads). Declared here so the allocating280
/// fronts below can call them; defined + explicitly instantiated in routines.cpp.281
/**282
* Outer product into the caller's buffer @p out (out FIRST) — the buffer-reuse overload of283
* @ref outer, writing the rank-1 result straight into @p out with no allocation.284
* @param out destination; a contiguous n×m matrix, overwritten. Must NOT alias @p a or @p b.285
* @param a length-n vector.286
* @param b length-m vector.287
* @complexity O(n·m).288
* @alloc none for contiguous operands (result written straight into @p out); a289
* non-contiguous operand is packed once into scratch.290
* @test LinalgRoutines.OuterIntoReusesBuffer291
*/292
template <ndarray::Field T, template <typename> class Array>293
requires HostArray<Array<T>>294
void outer(Array<T>& out, const Array<T>& a, const Array<T>& b);295
/**296
* Conjugate transpose into the caller's buffer @p out (out FIRST) — the buffer-reuse overload of297
* @ref conj_transpose, writing the c×r adjoint straight into @p out with no allocation.298
* @param out destination; a contiguous c×r matrix (for an r×c input), overwritten. Must NOT299
* alias @p a (it reads A while writing the transpose — not an in-place op).300
* @param a a 2-D matrix.301
* @complexity O(r·c).302
* @alloc none for a contiguous operand (adjoint written straight into @p out); a303
* non-contiguous operand is packed once into scratch.304
* @test LinalgRoutines.ConjTransposeIntoReusesBuffer305
*/306
template <ndarray::Field T, template <typename> class Array>307
requires HostArray<Array<T>>308
void conj_transpose(Array<T>& out, const Array<T>& a);309
/**310
* Kronecker product into the caller's buffer @p out (out FIRST) — the buffer-reuse overload of311
* @ref kron, writing the block product straight into @p out with no allocation.312
* @param out destination; a contiguous (m·p)×(k·q) matrix, overwritten. Must NOT alias @p a or @p b.313
* @param a m×k matrix.314
* @param b p×q matrix.315
* @complexity O(n⁴) in the output area.316
* @alloc none for contiguous operands (block product written straight into @p out); a317
* non-contiguous operand is packed once into scratch.318
* @test LinalgRoutines.KronIntoReusesBuffer319
*/320
template <ndarray::Field T, template <typename> class Array>321
requires HostArray<Array<T>>322
void kron(Array<T>& out, const Array<T>& a, const Array<T>& b);323
/// @endcond325
/**326
* Outer product of two vectors.327
*328
* Flattens both operands to vectors and forms the full rank-1 matrix; any pair329
* of vector lengths is accepted (no matching constraint). Allocates the result via330
* `Array<T>::uninitialized` and fills it through the out-parameter kernel — the host SIMD331
* path, or a device shader when `Array` is a device container (selected by concept).332
* @param a length-n vector.333
* @param b length-m vector.334
* @return n×m matrix aᵢbⱼ.335
* @complexity O(n·m).336
* @alloc allocates only the n×m result; operands read in place when contiguous.337
* @test LinalgRoutines.VdotInnerOuterKron338
* @crtest LinalgCompileRun.Outer339
* @systest StdlibE2E.Linalg340
*/341
template <ndarray::Field T, template <typename> class Array>342
requires NumericArray<Array<T>>343
[[nodiscard]] Array<T> outer(const Array<T>& a, const Array<T>& b) {344
Array<T> out = Array<T>::uninitialized({vector_len(a), vector_len(b)});345
outer(out, a, b);346
return out;347
}349
/**350
* Conjugate transpose (Hermitian adjoint) Aᴴ: transpose, then conjugate every entry (a plain351
* transpose for a real element — the conjugation is compiled out). A matrix is Hermitian iff352
* `conj_transpose(A) == A`.353
* @param a a 2-D matrix.354
* @return the c×r adjoint of an r×c input; throws on non-2-D input.355
* @complexity O(r·c).356
* @alloc allocates only the c×r result; a non-contiguous operand is packed once into scratch.357
* @test LinalgRoutines.ComplexProducts358
* @crtest LinalgCompileRun.ConjTranspose359
* @systest StdlibE2E.LinalgComplex360
*/361
template <ndarray::Field T, template <typename> class Array>362
requires NumericArray<Array<T>>363
[[nodiscard]] Array<T> conj_transpose(const Array<T>& a) {364
if (a.ndim() != 2) throw std::runtime_error("linalg: expected a 2-D matrix");365
Array<T> out = Array<T>::uninitialized({a.shape()[1], a.shape()[0]});366
conj_transpose(out, a);367
return out;368
}370
/**371
* Kronecker product.372
*373
* Requires both operands to be 2-D (throws otherwise) and replaces each entry of374
* @p a with that scalar times the whole of @p b, giving the (m·p)×(k·q) block375
* matrix; no dimension matching is needed.376
* @param a m×k matrix.377
* @param b p×q matrix.378
* @return (m·p)×(k·q) block product.379
* @complexity O(n⁴) in the output area.380
* @alloc allocates only the (m·p)×(k·q) result; a non-contiguous operand packs once into scratch.381
* @test LinalgRoutines.VdotInnerOuterKron382
* @crtest LinalgCompileRun.Kron383
* @systest StdlibE2E.Linalg384
*/385
template <ndarray::Field T, template <typename> class Array>386
requires NumericArray<Array<T>>387
[[nodiscard]] Array<T> kron(const Array<T>& a, const Array<T>& b) {388
if (a.ndim() != 2 || b.ndim() != 2)389
throw std::runtime_error("linalg: kron expects 2-D matrices");390
// Each output dimension is a PRODUCT of two input dims, so it must be overflow-checked391
// BEFORE it collapses into the shape vector: `product({m*p, k*q})` would only catch a wrap392
// of the final area, not a wrap of m*p (or k*q) alone, which under-allocates and lets the393
// kernel write out of bounds. `product({x, y})` is the shared checked multiply — it throws394
// on wrap instead. (See ndarray::detail::product; same class as the ndarray shape-overflow fix.)395
const std::size_t orows = ndarray::detail::product({a.shape()[0], b.shape()[0]});396
const std::size_t ocols = ndarray::detail::product({a.shape()[1], b.shape()[1]});397
Array<T> out = Array<T>::uninitialized({orows, ocols});398
kron(out, a, b);399
return out;400
}402
} // namespace cheatah::linalg