cheatah
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 once
5/**
6 * @file backend.hpp
7 * @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 template
10 * `Array`, with a `requires` concept enforcing the container. Both operands are spelled
11 * `Array<T>`, so a host⊗device or f64⊗f32 mix cannot deduce a single `Array`/`T` and is a
12 * compile error — the location/element firewall is FREE, via deduction, with no runtime
13 * 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>&)` that
17 * 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: the
19 * 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 the
22 * extension. (No CPO objects, no tag_invoke — plain concept-constrained overloads.)
23 */
24#include <stdexcept>
25#include <vector>
27#include "concepts.hpp"
29namespace cheatah::linalg {
31/// @cond INTERNAL
32/// Shared BATCHED shape validation. Both the allocating front and the out-parameter kernel run
33/// 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-element
36/// shape vector before any check ran.
37template <ndarray::Field T, template <typename> class Array>
38inline 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");
46/// @endcond
48/// @cond INTERNAL
49/// the allocation-free out-parameter kernel (HOST overload; a device
50/// extension adds its own `requires DeviceArray<Array<T>>` overload). Declared here so the
51/// 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 hands
54 * 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); a
61 * non-contiguous operand is packed once into scratch.
62 * @test LinalgRoutines.MatmulIntoReusesBuffer
63 * @test LinalgRoutines.ComplexMatmulIntoReusesBuffer
64 */
65template <ndarray::Field T, template <typename> class Array>
66 requires HostArray<Array<T>>
67void matmul(Array<T>& out, const Array<T>& a, const Array<T>& b);
68/// @endcond
70/**
71 * Matrix multiply — the allocating front. Both operands are `Array<T>` (so host⊗device / element
72 * mixes fail to deduce and are compile errors); requires both to be 2-D with matching inner
73 * dimensions — or both 3-D for the BATCHED product `[B,M,K] @ [B,K,N] → [B,M,N]` (equal batch
74 * counts, strict: no broadcast batching). Allocates the result via `Array<T>::uninitialized` and
75 * fills it through the out-parameter kernel — the host SIMD path, or a device shader when
76 * `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); parallelize
85 * across independent products in the caller.
86 * @test LinalgRoutines.ProductsAndTrace
87 * @test LinalgRoutines.BatchedMatmul
88 * @crtest LinalgCompileRun.Matmul
89 * @systest StdlibE2E.Linalg
90 */
91template <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;
109/**
110 * The flattened length of a vector-shaped operand — 1-D, or 2-D with a size-1 row/column
111 * (throws otherwise). Reads only host-resident shape metadata, so it is valid for ANY located
112 * 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.ProductsAndTrace
119 */
120template <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");
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 the
129// unqualified `op(out, …)`, which resolves to the HOST kernel below (routines.cpp) or a device
130// extension's `requires DeviceArray<Array<T>>` overload via ADL.
132/// @cond INTERNAL
133/// the scalar-out reduction kernels (HOST overloads; a device extension adds its
134/// own `requires DeviceArray<Array<T>>` overloads). Declared here so the allocating fronts below
135/// 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 analogue
138 * 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.ProductsAndTrace
142 */
143template <ndarray::Field T, template <typename> class Array>
144 requires HostArray<Array<T>>
145void 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 — the
148 * 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.VdotInnerOuterKron
152 */
153template <ndarray::Field T, template <typename> class Array>
154 requires HostArray<Array<T>>
155void 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 for
158 * 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.VdotInnerOuterKron
162 */
163template <ndarray::Field T, template <typename> class Array>
164 requires HostArray<Array<T>>
165void 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.ProductsAndTrace
171 */
172template <ndarray::Field T, template <typename> class Array>
173 requires HostArray<Array<T>>
174void trace(T& out, const Array<T>& a);
175/// @endcond
177/**
178 * Dot product: 1-D inner product (vectors flattened) — the bilinear Σ aᵢbᵢ. Flattens each
179 * operand to a vector (1-D, or 2-D with a size-1 row/column) and throws if either is not
180 * 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.ProductsAndTrace
188 * @test LinalgRoutines.ComplexProducts
189 * @crtest LinalgCompileRun.Dot
190 * @crtest LinalgCompileRun.ComplexDot
191 * @systest StdlibE2E.Linalg
192 * @systest StdlibE2E.LinalgComplex
193 */
194template <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;
204/**
205 * Vector dot product. For a REAL element this is the bilinear Σ aᵢbᵢ (identical to @ref dot and
206 * @ref inner); for a **complex** element it is the conjugate-linear Hermitian inner product
207 * ⟨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.VdotInnerOuterKron
215 * @test LinalgRoutines.ComplexProducts
216 * @crtest LinalgCompileRun.Vdot
217 * @crtest LinalgCompileRun.ComplexVdot
218 * @systest StdlibE2E.Linalg
219 * @systest StdlibE2E.LinalgComplex
220 */
221template <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;
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.VdotInnerOuterKron
240 * @crtest LinalgCompileRun.Inner
241 * @systest StdlibE2E.Linalg
242 */
243template <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;
253/**
254 * Trace: the sum of the matrix diagonal, as the scalar `T`. Requires a 2-D matrix (throws
255 * 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.ProductsAndTrace
263 * @crtest LinalgCompileRun.Trace
264 * @systest StdlibE2E.Linalg
265 */
266template <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;
275// ---- products with array results (outer / conj_transpose / kron): the matmul pattern ----
277/// @cond INTERNAL
278/// the allocation-free out-parameter kernels (HOST overloads; a device extension
279/// adds its own `requires DeviceArray<Array<T>>` overloads). Declared here so the allocating
280/// 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 of
283 * @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); a
289 * non-contiguous operand is packed once into scratch.
290 * @test LinalgRoutines.OuterIntoReusesBuffer
291 */
292template <ndarray::Field T, template <typename> class Array>
293 requires HostArray<Array<T>>
294void 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 of
297 * @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 NOT
299 * 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); a
303 * non-contiguous operand is packed once into scratch.
304 * @test LinalgRoutines.ConjTransposeIntoReusesBuffer
305 */
306template <ndarray::Field T, template <typename> class Array>
307 requires HostArray<Array<T>>
308void 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 of
311 * @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); a
317 * non-contiguous operand is packed once into scratch.
318 * @test LinalgRoutines.KronIntoReusesBuffer
319 */
320template <ndarray::Field T, template <typename> class Array>
321 requires HostArray<Array<T>>
322void kron(Array<T>& out, const Array<T>& a, const Array<T>& b);
323/// @endcond
325/**
326 * Outer product of two vectors.
327 *
328 * Flattens both operands to vectors and forms the full rank-1 matrix; any pair
329 * of vector lengths is accepted (no matching constraint). Allocates the result via
330 * `Array<T>::uninitialized` and fills it through the out-parameter kernel — the host SIMD
331 * 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.VdotInnerOuterKron
338 * @crtest LinalgCompileRun.Outer
339 * @systest StdlibE2E.Linalg
340 */
341template <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;
349/**
350 * Conjugate transpose (Hermitian adjoint) Aᴴ: transpose, then conjugate every entry (a plain
351 * transpose for a real element — the conjugation is compiled out). A matrix is Hermitian iff
352 * `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.ComplexProducts
358 * @crtest LinalgCompileRun.ConjTranspose
359 * @systest StdlibE2E.LinalgComplex
360 */
361template <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;
370/**
371 * Kronecker product.
372 *
373 * Requires both operands to be 2-D (throws otherwise) and replaces each entry of
374 * @p a with that scalar times the whole of @p b, giving the (m·p)×(k·q) block
375 * 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.VdotInnerOuterKron
382 * @crtest LinalgCompileRun.Kron
383 * @systest StdlibE2E.Linalg
384 */
385template <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-checked
391 // BEFORE it collapses into the shape vector: `product({m*p, k*q})` would only catch a wrap
392 // of the final area, not a wrap of m*p (or k*q) alone, which under-allocates and lets the
393 // kernel write out of bounds. `product({x, y})` is the shared checked multiply — it throws
394 // 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;
402} // namespace cheatah::linalg