Source
stdlib/tests/linalg_routines_test.cpp
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
#include "ndarray.hpp"4
#include "routines.hpp"6
#include <cstddef>7
#include <algorithm>8
#include <cmath>9
#include <complex>10
#include <stdexcept>11
#include <vector>13
#include <gtest/gtest.h>15
namespace nd = cheatah::ndarray;16
namespace la = cheatah::linalg;18
namespace {19
nd::NDArray mat(std::size_t r, std::size_t c, std::vector<double> data) {20
return nd::reshape(nd::array(std::move(data)), {(long long)r, (long long)c});21
}22
bool close(double a, double b, double tol = 1e-9) { return std::fabs(a - b) < tol; }23
// The general eig()/eigvals() return a complex spectrum; these read one element and24
// compare against a complex (or, implicitly, a real) expectation.25
std::complex<double> cget(const la::CNDArray& v, const std::vector<long long>& idx) {26
return nd::get(v, idx);27
}28
bool cclose(std::complex<double> a, std::complex<double> b, double tol = 1e-6) {29
return std::abs(a - b) < tol;30
}31
using C = std::complex<double>;32
la::CNDArray cvec(std::vector<C> data) { return nd::array(std::move(data)); }33
la::CNDArray cmat(std::size_t r, std::size_t c, std::vector<C> data) {34
return nd::reshape(nd::array(std::move(data)), {(long long)r, (long long)c});35
}36
} // namespace38
TEST(LinalgRoutines, ProductsAndTrace) {39
const nd::NDArray a = mat(2, 3, {1, 2, 3, 4, 5, 6});40
const nd::NDArray b = mat(3, 2, {7, 8, 9, 10, 11, 12});41
const nd::NDArray c = la::matmul(a, b); // [[58,64],[139,154]]42
EXPECT_DOUBLE_EQ(nd::get(c, {0, 0}), 58);43
EXPECT_DOUBLE_EQ(nd::get(c, {0, 1}), 64);44
EXPECT_DOUBLE_EQ(nd::get(c, {1, 1}), 154);45
EXPECT_DOUBLE_EQ(la::dot(nd::array({1.0, 2.0, 3.0}), nd::array({4.0, 5.0, 6.0})), 32); // 4+10+1846
EXPECT_DOUBLE_EQ(la::trace(mat(2, 2, {1, 2, 3, 4})), 5); // 1+447
}49
// The user-provided-output overload writes the SAME result into the caller's buffer with NO50
// reallocation: the buffer's data pointer is identical before and after, and the values are correct.51
TEST(LinalgRoutines, MatmulIntoReusesBuffer) {52
const nd::NDArray a = mat(2, 3, {1, 2, 3, 4, 5, 6});53
const nd::NDArray b = mat(3, 2, {7, 8, 9, 10, 11, 12});54
nd::NDArray out = nd::zeros({2, 2});55
const double* const before = out.buffer()->data() + out.offset(); // capture the buffer identity56
la::matmul(out, a, b); // [[58,64],[139,154]] into out57
EXPECT_EQ(out.buffer()->data() + out.offset(), before); // SAME buffer — no reallocation58
EXPECT_DOUBLE_EQ(nd::get(out, {0, 0}), 58);59
EXPECT_DOUBLE_EQ(nd::get(out, {0, 1}), 64);60
EXPECT_DOUBLE_EQ(nd::get(out, {1, 0}), 139);61
EXPECT_DOUBLE_EQ(nd::get(out, {1, 1}), 154);62
// matches the allocating overload exactly63
const nd::NDArray c = la::matmul(a, b);64
EXPECT_DOUBLE_EQ(nd::get(out, {1, 1}), nd::get(c, {1, 1}));65
// a wrong-shaped out, and an out that aliases an input (matmul is not in-place), are rejected.66
nd::NDArray wrong = nd::zeros({3, 3});67
EXPECT_THROW(static_cast<void>(la::matmul(wrong, a, b)), std::runtime_error);68
nd::NDArray sq = mat(2, 2, {1, 2, 3, 4});69
EXPECT_THROW(static_cast<void>(la::matmul(sq, sq, sq)), std::runtime_error);70
// a non-2-D operand to the out-form is rejected up front ("expects 2-D matrices").71
nd::NDArray vec = nd::array({1.0, 2.0, 3.0}); // 1-D72
nd::NDArray out2 = nd::zeros({2, 2});73
EXPECT_THROW(static_cast<void>(la::matmul(out2, vec, b)), std::runtime_error);74
// A 3-D `a` against a 2-D `b` took the BATCHED branch and read b.shape()[2] — past the end of75
// a two-element shape vector — because the rank checks lived only in the allocating front.76
// The out-form is public, so it validates for itself.77
nd::NDArray batch = nd::zeros({2, 2, 3});78
nd::NDArray out3 = nd::zeros({2, 2, 2});79
EXPECT_THROW(static_cast<void>(la::matmul(out3, batch, b)), std::runtime_error);80
// and the batched checks the front used to own are enforced here too81
nd::NDArray batch_b = nd::zeros({3, 3, 2}); // mismatched batch count82
EXPECT_THROW(static_cast<void>(la::matmul(out3, batch, batch_b)), std::runtime_error);83
nd::NDArray batch_k = nd::zeros({2, 4, 2}); // mismatched contracted dimension84
EXPECT_THROW(static_cast<void>(la::matmul(out3, batch, batch_k)), std::runtime_error);85
}87
// Every product/least-squares front validates its operand dimensions and throws on a mismatch —88
// the error paths that a happy-path test never reaches. dot/vdot/inner reject unequal vector89
// lengths; batched matmul (two 3-D operands) rejects a mismatched contracted dimension; lstsq90
// rejects a row-count mismatch.91
TEST(LinalgRoutines, DimensionMismatchThrows) {92
const nd::NDArray v2 = nd::array({1.0, 2.0});93
const nd::NDArray v3 = nd::array({1.0, 2.0, 3.0});94
EXPECT_THROW(static_cast<void>(la::dot(v2, v3)), std::runtime_error);95
EXPECT_THROW(static_cast<void>(la::vdot(v2, v3)), std::runtime_error);96
EXPECT_THROW(static_cast<void>(la::inner(v2, v3)), std::runtime_error);97
// Batched [B,M,K] @ [B,K,N]: equal batch counts but a mismatched inner dim (K vs K').98
const nd::NDArray a3 = nd::reshape(nd::array(std::vector<double>(std::size_t{2} * 3 * 4, 1.0)), {2, 3, 4});99
const nd::NDArray b3 = nd::reshape(nd::array(std::vector<double>(std::size_t{2} * 5 * 6, 1.0)), {2, 5, 6});100
EXPECT_THROW(static_cast<void>(la::matmul(a3, b3)), std::runtime_error); // K=4 != K'=5101
// lstsq: A (m×n) and b (m×k) must share the row count m.102
const nd::NDArray A = mat(3, 2, {1, 2, 3, 4, 5, 6});103
const nd::NDArray rhs = mat(2, 1, {1, 2}); // 2 rows != A's 3104
EXPECT_THROW(static_cast<void>(la::lstsq(A, rhs)), std::runtime_error);105
}107
// The memory-bound products / transposes have GENUINELY zero-allocation out-param overloads: they108
// write their kernel straight into the caller's buffer (data pointer identical before/after), match109
// the allocating overload, and reject a wrong shape or an out that aliases an input.110
TEST(LinalgRoutines, OuterIntoReusesBuffer) {111
const nd::NDArray a = nd::array({1.0, 2.0, 3.0});112
const nd::NDArray b = nd::array({4.0, 5.0});113
nd::NDArray out = nd::zeros({3, 2});114
const auto* before = out.buffer().get();115
la::outer(out, a, b);116
EXPECT_EQ(out.buffer().get(), before); // SAME buffer — no reallocation117
const nd::NDArray ref = la::outer(a, b);118
EXPECT_DOUBLE_EQ(nd::get(out, {0, 0}), 4.0); // 1*4119
EXPECT_DOUBLE_EQ(nd::get(out, {2, 1}), nd::get(ref, {2, 1})); // 3*5, matches allocating form120
nd::NDArray wrong = nd::zeros({2, 2});121
EXPECT_THROW(static_cast<void>(la::outer(wrong, a, b)), std::runtime_error);122
}124
TEST(LinalgRoutines, KronIntoReusesBuffer) {125
const nd::NDArray a = mat(2, 2, {1, 0, 0, 1});126
const nd::NDArray b = mat(2, 2, {1, 2, 3, 4});127
nd::NDArray out = nd::zeros({4, 4});128
const auto* before = out.buffer().get();129
la::kron(out, a, b);130
EXPECT_EQ(out.buffer().get(), before);131
const nd::NDArray ref = la::kron(a, b);132
EXPECT_DOUBLE_EQ(nd::get(out, {0, 1}), nd::get(ref, {0, 1}));133
EXPECT_DOUBLE_EQ(nd::get(out, {3, 3}), nd::get(ref, {3, 3}));134
// A non-2-D operand is rejected up front ("kron expects 2-D matrices").135
nd::NDArray vec = nd::array({1.0, 2.0}); // 1-D136
EXPECT_THROW(static_cast<void>(la::kron(out, vec, b)), std::runtime_error);137
// An out that ALIASES an input is rejected by reject_alias (kron is not computed in place).138
nd::NDArray alias = mat(2, 2, {1, 0, 0, 1});139
EXPECT_THROW(static_cast<void>(la::kron(alias, alias, b)), std::runtime_error);140
}142
TEST(LinalgRoutines, ConjTransposeIntoReusesBuffer) {143
const la::CNDArray M = cmat(2, 3, {C(1, 1), C(2, 0), C(3, -1), C(0, 2), C(1, 0), C(4, 4)});144
la::CNDArray out = cmat(3, 2, std::vector<C>(6));145
const auto* before = out.buffer().get();146
la::conj_transpose(out, M);147
EXPECT_EQ(out.buffer().get(), before);148
const la::CNDArray ref = la::conj_transpose(M);149
EXPECT_TRUE(cclose(cget(out, {0, 0}), C(1, -1)));150
EXPECT_TRUE(cclose(cget(out, {2, 1}), cget(ref, {2, 1})));151
}153
TEST(LinalgRoutines, ComplexMatmulIntoReusesBuffer) {154
const la::CNDArray a = cmat(2, 3, {C(1, 0), C(2, 0), C(3, 0), C(4, 0), C(5, 0), C(6, 0)});155
const la::CNDArray b = cmat(3, 2, {C(1, 1), C(0, 0), C(0, 1), C(1, 0), C(2, 0), C(0, 1)});156
la::CNDArray out = cmat(2, 2, std::vector<C>(4));157
const auto* before = out.buffer().get();158
la::matmul(out, a, b);159
EXPECT_EQ(out.buffer().get(), before);160
const la::CNDArray ref = la::matmul(a, b);161
EXPECT_TRUE(cclose(cget(out, {0, 0}), cget(ref, {0, 0})));162
EXPECT_TRUE(cclose(cget(out, {1, 1}), cget(ref, {1, 1})));163
}165
// The O(n³) factorizations reuse the caller's OUTPUT buffer (data pointer identical before/after)166
// and match the allocating overload — their internal factorization workspace is allocated regardless.167
TEST(LinalgRoutines, FactorizationOutReusesBuffer) {168
const nd::NDArray A = mat(2, 2, {4, 3, 6, 3});169
const nd::NDArray spd = mat(2, 2, {4, 2, 2, 3});170
const nd::NDArray sym = mat(2, 2, {2, 1, 1, 2});171
const nd::NDArray gen = mat(2, 2, {2, 0, 0, 5});172
{ // solve173
nd::NDArray out = nd::zeros({2});174
const auto* b = out.buffer().get();175
la::solve(out, A, nd::array({10.0, 12.0}));176
EXPECT_EQ(out.buffer().get(), b);177
EXPECT_TRUE(close(nd::get(out, {0}), 1.0));178
EXPECT_TRUE(close(nd::get(out, {1}), 2.0));179
}180
{ // inv181
nd::NDArray out = nd::zeros({2, 2});182
const auto* b = out.buffer().get();183
la::inv(out, A);184
EXPECT_EQ(out.buffer().get(), b);185
const nd::NDArray ref = la::inv(A);186
EXPECT_TRUE(close(nd::get(out, {0, 0}), nd::get(ref, {0, 0})));187
}188
{ // lstsq (2-D column rhs); on a square system it equals solve189
nd::NDArray out = nd::zeros({2, 1});190
const auto* b = out.buffer().get();191
la::lstsq(out, A, mat(2, 1, {10, 12}));192
EXPECT_EQ(out.buffer().get(), b);193
EXPECT_TRUE(close(nd::get(out, {0, 0}), 1.0, 1e-6));194
}195
{ // cholesky196
nd::NDArray out = nd::zeros({2, 2});197
const auto* b = out.buffer().get();198
la::cholesky(out, spd);199
EXPECT_EQ(out.buffer().get(), b);200
const nd::NDArray ref = la::cholesky(spd);201
EXPECT_TRUE(close(nd::get(out, {0, 0}), nd::get(ref, {0, 0})));202
}203
{ // pinv204
nd::NDArray out = nd::zeros({2, 2});205
const auto* b = out.buffer().get();206
la::pinv(out, A);207
EXPECT_EQ(out.buffer().get(), b);208
const nd::NDArray ref = la::pinv(A);209
EXPECT_TRUE(close(nd::get(out, {0, 0}), nd::get(ref, {0, 0}), 1e-6));210
}211
{ // matrix_power212
nd::NDArray out = nd::zeros({2, 2});213
const auto* b = out.buffer().get();214
la::matrix_power(out, A, 2);215
EXPECT_EQ(out.buffer().get(), b);216
const nd::NDArray ref = la::matrix_power(A, 2);217
EXPECT_TRUE(close(nd::get(out, {0, 0}), nd::get(ref, {0, 0})));218
}219
{ // svdvals220
nd::NDArray out = nd::zeros({2});221
const auto* b = out.buffer().get();222
la::svdvals(out, sym);223
EXPECT_EQ(out.buffer().get(), b);224
EXPECT_TRUE(close(nd::get(out, {0}), 3.0, 1e-6));225
}226
{ // eigvalsh (symmetric)227
nd::NDArray out = nd::zeros({2});228
const auto* b = out.buffer().get();229
la::eigvalsh(out, sym);230
EXPECT_EQ(out.buffer().get(), b);231
EXPECT_TRUE(close(nd::get(out, {0}), 3.0, 1e-6));232
}233
{ // eigvals (general, complex out)234
la::CNDArray out = cvec(std::vector<C>(2));235
const auto* b = out.buffer().get();236
la::eigvals(out, gen);237
EXPECT_EQ(out.buffer().get(), b);238
EXPECT_TRUE(cclose(cget(out, {0}), 5.0));239
}240
{ // eigvalsh (complex Hermitian, real out)241
const la::CNDArray H = cmat(2, 2, {C(2, 0), C(1, 1), C(1, -1), C(3, 0)});242
nd::NDArray out = nd::zeros({2});243
const auto* b = out.buffer().get();244
la::eigvalsh(out, H);245
EXPECT_EQ(out.buffer().get(), b);246
EXPECT_TRUE(close(nd::get(out, {0}), 4.0, 1e-6));247
}248
}250
// The multi-output decompositions reuse EVERY caller-provided output buffer (one per factor).251
TEST(LinalgRoutines, DecompositionOutReusesBuffer) {252
{ // qr253
const nd::NDArray A = mat(3, 2, {1, 0, 1, 1, 0, 1});254
nd::NDArray q = nd::zeros({3, 2}), r = nd::zeros({2, 2});255
const auto* bq = q.buffer().get();256
const auto* br = r.buffer().get();257
la::qr(q, r, A);258
EXPECT_EQ(q.buffer().get(), bq);259
EXPECT_EQ(r.buffer().get(), br);260
const la::QR ref = la::qr(A);261
EXPECT_TRUE(close(nd::get(q, {0, 0}), nd::get(ref.q, {0, 0}), 1e-6));262
EXPECT_TRUE(close(nd::get(r, {0, 0}), nd::get(ref.r, {0, 0}), 1e-6));263
}264
{ // svd265
const nd::NDArray A = mat(2, 2, {2, 0, 0, 3});266
nd::NDArray u = nd::zeros({2, 2}), s = nd::zeros({2}), vh = nd::zeros({2, 2});267
const auto* bu = u.buffer().get();268
const auto* bs = s.buffer().get();269
const auto* bv = vh.buffer().get();270
la::svd(u, s, vh, A);271
EXPECT_EQ(u.buffer().get(), bu);272
EXPECT_EQ(s.buffer().get(), bs);273
EXPECT_EQ(vh.buffer().get(), bv);274
EXPECT_TRUE(close(nd::get(s, {0}), 3.0, 1e-6));275
}276
{ // eigh (symmetric, real)277
nd::NDArray vals = nd::zeros({2}), vecs = nd::zeros({2, 2});278
const auto* bvl = vals.buffer().get();279
const auto* bvc = vecs.buffer().get();280
la::eigh(vals, vecs, mat(2, 2, {2, 1, 1, 2}));281
EXPECT_EQ(vals.buffer().get(), bvl);282
EXPECT_EQ(vecs.buffer().get(), bvc);283
EXPECT_TRUE(close(nd::get(vals, {0}), 3.0, 1e-6));284
}285
{ // eigh (complex Hermitian: real values, complex vectors)286
const la::CNDArray H = cmat(2, 2, {C(2, 0), C(1, 1), C(1, -1), C(3, 0)});287
nd::NDArray vals = nd::zeros({2});288
la::CNDArray vecs = cmat(2, 2, std::vector<C>(4));289
const auto* bvl = vals.buffer().get();290
const auto* bvc = vecs.buffer().get();291
la::eigh(vals, vecs, H);292
EXPECT_EQ(vals.buffer().get(), bvl);293
EXPECT_EQ(vecs.buffer().get(), bvc);294
EXPECT_TRUE(close(nd::get(vals, {0}), 4.0, 1e-6));295
}296
{ // eig (general, complex values + vectors)297
la::CNDArray vals = cvec(std::vector<C>(2));298
la::CNDArray vecs = cmat(2, 2, std::vector<C>(4));299
const auto* bvl = vals.buffer().get();300
const auto* bvc = vecs.buffer().get();301
la::eig(vals, vecs, mat(2, 2, {2, 0, 0, 5}));302
EXPECT_EQ(vals.buffer().get(), bvl);303
EXPECT_EQ(vecs.buffer().get(), bvc);304
EXPECT_TRUE(cclose(cget(vals, {0}), 5.0));305
}306
}308
TEST(LinalgRoutines, SolveDetInv) {309
const nd::NDArray A = mat(2, 2, {4, 3, 6, 3}); // det = 12-18 = -6310
EXPECT_TRUE(close(la::det(A), -6.0));311
const nd::NDArray x = la::solve(A, nd::array({10.0, 12.0})); // 4x+3y=10, 6x+3y=12 -> x=1,y=2312
EXPECT_TRUE(close(nd::get(x, {0}), 1.0));313
EXPECT_TRUE(close(nd::get(x, {1}), 2.0));314
const nd::NDArray Ai = la::inv(A);315
const nd::NDArray I = la::matmul(A, Ai); // identity316
EXPECT_TRUE(close(nd::get(I, {0, 0}), 1.0));317
EXPECT_TRUE(close(nd::get(I, {0, 1}), 0.0));318
EXPECT_TRUE(close(nd::get(I, {1, 1}), 1.0));319
}321
TEST(LinalgRoutines, CholeskyAndQR) {322
const nd::NDArray A = mat(2, 2, {4, 2, 2, 3}); // SPD323
const nd::NDArray L = la::cholesky(A);324
const nd::NDArray LLt = la::matmul(L, la::matmul(la::inv(L), A)); // == A trivially; check L Lᵀ:325
// verify L·Lᵀ == A326
const nd::NDArray Lt = mat(2, 2, {nd::get(L, {0, 0}), nd::get(L, {1, 0}),327
nd::get(L, {0, 1}), nd::get(L, {1, 1})});328
const nd::NDArray rec = la::matmul(L, Lt);329
EXPECT_TRUE(close(nd::get(rec, {0, 0}), 4.0));330
EXPECT_TRUE(close(nd::get(rec, {1, 1}), 3.0));332
const la::QR qr = la::qr(mat(3, 2, {1, 0, 1, 1, 0, 1}));333
const nd::NDArray QtQ = la::matmul(la::pinv(qr.q), qr.q); // Q has orthonormal cols334
EXPECT_TRUE(close(nd::get(QtQ, {0, 0}), 1.0, 1e-6));335
const nd::NDArray reQR = la::matmul(qr.q, qr.r); // == original336
EXPECT_TRUE(close(nd::get(reQR, {0, 0}), 1.0, 1e-6));337
EXPECT_TRUE(close(nd::get(reQR, {1, 1}), 1.0, 1e-6));338
}340
TEST(LinalgRoutines, SvdAndEigh) {341
const nd::NDArray A = mat(2, 2, {2, 0, 0, 3});342
const la::SVD s = la::svd(A);343
EXPECT_TRUE(close(nd::get(s.s, {0}), 3.0, 1e-6)); // singular values 3, 2 (descending)344
EXPECT_TRUE(close(nd::get(s.s, {1}), 2.0, 1e-6));345
// svdvals (values-only fast path) agrees with svd().s346
const nd::NDArray sv = la::svdvals(A);347
EXPECT_TRUE(close(nd::get(sv, {0}), 3.0, 1e-6));348
EXPECT_TRUE(close(nd::get(sv, {1}), 2.0, 1e-6));350
// symmetric eigen: [[2,1],[1,2]] -> eigenvalues 3, 1351
const la::Eig e = la::eigh(mat(2, 2, {2, 1, 1, 2}));352
EXPECT_TRUE(close(nd::get(e.values, {0}), 3.0, 1e-6));353
EXPECT_TRUE(close(nd::get(e.values, {1}), 1.0, 1e-6));355
// general eigenvalues of [[2,0],[0,5]] -> 5, 2 (real, returned as complex)356
const la::CNDArray ev = la::eigvals(mat(2, 2, {2, 0, 0, 5}));357
EXPECT_TRUE(cclose(cget(ev, {0}), 5.0));358
EXPECT_TRUE(cclose(cget(ev, {1}), 2.0));359
}361
TEST(LinalgRoutines, ComplexProducts) {362
const la::CNDArray a = cvec({C(1, 2), C(3, -1)});363
const la::CNDArray b = cvec({C(0, 1), C(2, 0)});364
// Bilinear dot (no conjugation): (1+2j)(0+1j) + (3-1j)(2) = (-2+1j) + (6-2j) = 4-1j.365
EXPECT_TRUE(cclose(la::dot(a, b), C(4, -1)));366
// Hermitian inner product (conjugate the first): conj(a)·b = (1-2j)(0+1j)+(3+1j)(2) = (2+1j)+(6+2j) = 8+3j.367
EXPECT_TRUE(cclose(la::vdot(a, b), C(8, 3)));368
// vdot(a,a) is the real squared norm ‖a‖² = 1+4+9+1 = 15.369
EXPECT_TRUE(cclose(la::vdot(a, a), C(15, 0)));371
// Conjugate transpose (Hermitian adjoint): transpose + conjugate every entry.372
const la::CNDArray M = cmat(2, 2, {C(1, 1), C(2, 0), C(0, 0), C(3, -1)});373
const la::CNDArray H = la::conj_transpose(M); // [[1-1j, 0],[2, 3+1j]]374
EXPECT_TRUE(cclose(cget(H, {0, 0}), C(1, -1)));375
EXPECT_TRUE(cclose(cget(H, {0, 1}), C(0, 0)));376
EXPECT_TRUE(cclose(cget(H, {1, 0}), C(2, 0)));377
EXPECT_TRUE(cclose(cget(H, {1, 1}), C(3, 1)));379
// Complex matmul: M · Mᴴ is Hermitian; check entry (0,0) = |1+1j|² + |2|² = 2 + 4 = 6.380
const la::CNDArray P = la::matmul(M, H);381
EXPECT_TRUE(cclose(cget(P, {0, 0}), C(6, 0)));382
EXPECT_TRUE(cclose(cget(P, {1, 1}), C(10, 0))); // |0|² + |3-1j|² = 0 + 10384
// as_cvector accepts a 2-D N×1 / 1×N as a flat vector (like the real path)…385
const la::CNDArray col = cmat(2, 1, {C(1, 0), C(0, 1)});386
EXPECT_TRUE(cclose(la::dot(col, col), C(0, 0))); // 1·1 + i·i = 1 − 1 = 0387
// …and rejects a genuine 2-D matrix where a vector is required.388
EXPECT_THROW(static_cast<void>(la::vdot(M, M)), std::runtime_error);389
}391
TEST(LinalgRoutines, GeneralEigVectors) {392
// For a real matrix, eig() returns complex eigenvalues AND eigenvectors (via393
// inverse iteration). Verify A·v_k = λ_k·v_k for each column (phase-independent),394
// building a complex copy Ac of A so we can multiply the complex eigenvectors.395
const auto check = [](const std::vector<double>& data, std::size_t n) {396
const nd::NDArray A = mat(n, n, data);397
std::vector<C> cdata;398
cdata.reserve(data.size());399
for (double x : data) cdata.emplace_back(x, 0.0);400
const la::CNDArray Ac = cmat(n, n, cdata);401
const la::EigC e = la::eig(A);402
const la::CNDArray AV = la::matmul(Ac, e.vectors);403
for (std::size_t k = 0; k < n; ++k) {404
const C lam = cget(e.values, {(long long)k});405
for (std::size_t r = 0; r < n; ++r) {406
EXPECT_TRUE(cclose(cget(AV, {(long long)r, (long long)k}),407
lam * cget(e.vectors, {(long long)r, (long long)k}), 1e-5));408
}409
}410
};411
check({2, 1, 0, 3}, 2); // real eigenvalues 3, 2 (upper-triangular)412
check({0, -1, 1, 0}, 2); // complex conjugate pair ±i (rotation)413
check({0, 1, 2, 0}, 2); // eigenvalues ±√2; forces a pivot row-swap in the solve414
check({2, 1, 1, 1, 2, 1, 1, 1, 2}, 3); // symmetric -> real eigenvalues 4,1,1415
}417
TEST(LinalgRoutines, ComplexHermitianEigh) {418
// H = [[2, 1+i],[1-i, 3]] is Hermitian (conj_transpose(H) == H); eigenvalues 4, 1.419
const la::CNDArray H = cmat(2, 2, {C(2, 0), C(1, 1), C(1, -1), C(3, 0)});420
const nd::NDArray w = la::eigvalsh(H); // real, descending421
EXPECT_TRUE(close(nd::get(w, {0}), 4.0, 1e-6));422
EXPECT_TRUE(close(nd::get(w, {1}), 1.0, 1e-6));424
const la::EighC e = la::eigh(H);425
EXPECT_TRUE(close(nd::get(e.values, {0}), 4.0, 1e-6));426
EXPECT_TRUE(close(nd::get(e.values, {1}), 1.0, 1e-6));427
// Verify the eigenpairs: H·V should equal V·diag(λ), independent of eigenvector428
// phase. So column k of H·V equals λ_k · column k of V.429
const la::CNDArray HV = la::matmul(H, e.vectors);430
for (int k = 0; k < 2; ++k) {431
const double lam = nd::get(e.values, {k});432
for (int r = 0; r < 2; ++r) {433
EXPECT_TRUE(cclose(cget(HV, {r, k}), lam * cget(e.vectors, {r, k})));434
}435
}436
// Eigenvectors are unit-norm: ⟨v,v⟩ = 1.437
for (int k = 0; k < 2; ++k) {438
const la::CNDArray vk = cvec({cget(e.vectors, {0, k}), cget(e.vectors, {1, k})});439
EXPECT_TRUE(cclose(la::vdot(vk, vk), C(1, 0)));440
}441
}443
TEST(LinalgRoutines, NormAndRank) {444
EXPECT_TRUE(close(la::norm(nd::array({3.0, 4.0})), 5.0)); // L2445
EXPECT_EQ(la::matrix_rank(mat(2, 2, {1, 2, 2, 4})), 1); // rank-deficient446
EXPECT_EQ(la::matrix_rank(mat(2, 2, {1, 0, 0, 1})), 2);447
}449
TEST(LinalgRoutines, VdotInnerOuterKron) {450
const nd::NDArray a = nd::array({1.0, 2.0, 3.0});451
const nd::NDArray b = nd::array({4.0, 5.0, 6.0});452
EXPECT_DOUBLE_EQ(la::vdot(a, b), 32.0); // 4+10+18453
EXPECT_DOUBLE_EQ(la::inner(a, b), 32.0);454
const nd::NDArray o = la::outer(nd::array({1.0, 2.0}), nd::array({3.0, 4.0})); // [[3,4],[6,8]]455
EXPECT_DOUBLE_EQ(nd::get(o, {0, 0}), 3.0);456
EXPECT_DOUBLE_EQ(nd::get(o, {1, 1}), 8.0);457
const nd::NDArray k = la::kron(mat(2, 2, {1, 0, 0, 1}), mat(2, 2, {1, 2, 3, 4})); // I⊗B458
EXPECT_DOUBLE_EQ(nd::get(k, {0, 0}), 1.0);459
EXPECT_DOUBLE_EQ(nd::get(k, {0, 1}), 2.0);460
EXPECT_DOUBLE_EQ(nd::get(k, {2, 2}), 1.0); // second diagonal block461
EXPECT_DOUBLE_EQ(nd::get(k, {3, 3}), 4.0);462
}464
TEST(LinalgRoutines, MatrixPower) {465
const nd::NDArray A = mat(2, 2, {2, 0, 0, 3});466
const nd::NDArray A0 = la::matrix_power(A, 0); // identity467
EXPECT_TRUE(close(nd::get(A0, {0, 0}), 1.0));468
EXPECT_TRUE(close(nd::get(A0, {1, 1}), 1.0));469
const nd::NDArray A3 = la::matrix_power(A, 3); // diag(8, 27)470
EXPECT_TRUE(close(nd::get(A3, {0, 0}), 8.0));471
EXPECT_TRUE(close(nd::get(A3, {1, 1}), 27.0));472
const nd::NDArray Am1 = la::matrix_power(A, -1); // diag(1/2, 1/3)473
EXPECT_TRUE(close(nd::get(Am1, {0, 0}), 0.5));474
EXPECT_TRUE(close(nd::get(Am1, {1, 1}), 1.0 / 3.0));475
}477
TEST(LinalgRoutines, SlogdetAndCond) {478
const la::SLogDet sd = la::slogdet(mat(2, 2, {4, 3, 6, 3})); // det = -6479
EXPECT_TRUE(close(sd.sign, -1.0));480
EXPECT_TRUE(close(sd.logabsdet, std::log(6.0), 1e-9));481
EXPECT_TRUE(close(la::cond(mat(2, 2, {2, 0, 0, 2})), 1.0, 1e-6)); // well-conditioned482
}484
TEST(LinalgRoutines, Lstsq) {485
// lstsq(a, b) = pinv(a) · b and takes a 2-D b (column vector). On a square486
// system, least-squares == solve.487
const nd::NDArray A = mat(2, 2, {4, 3, 6, 3});488
const nd::NDArray x = la::lstsq(A, mat(2, 1, {10, 12})); // -> [[1], [2]]489
EXPECT_TRUE(close(nd::get(x, {0, 0}), 1.0, 1e-6));490
EXPECT_TRUE(close(nd::get(x, {1, 0}), 2.0, 1e-6));491
}493
TEST(LinalgRoutines, EigvalshSymmetric) {494
const nd::NDArray w = la::eigvalsh(mat(2, 2, {2, 1, 1, 2})); // eigenvalues 1, 3495
const double v0 = nd::get(w, {0}), v1 = nd::get(w, {1});496
EXPECT_TRUE(close(std::min(v0, v1), 1.0, 1e-6));497
EXPECT_TRUE(close(std::max(v0, v1), 3.0, 1e-6));498
}500
TEST(LinalgRoutines, GeneralEig) {501
// Non-symmetric (upper-triangular) [[2,1],[0,3]] -> eigenvalues 2, 3 (real).502
const la::EigC e = la::eig(mat(2, 2, {2, 1, 0, 3}));503
EXPECT_TRUE(cclose(cget(e.values, {0}), 3.0)); // descending504
EXPECT_TRUE(cclose(cget(e.values, {1}), 2.0));505
// eigvals on the same non-symmetric matrix exercises the general path too.506
const la::CNDArray ev = la::eigvals(mat(2, 2, {2, 1, 0, 3}));507
EXPECT_TRUE(cclose(cget(ev, {0}), 3.0));508
EXPECT_TRUE(cclose(cget(ev, {1}), 2.0));509
}511
TEST(LinalgRoutines, GeneralEigOnSymmetricPromotesToComplex) {512
// eig() on a symmetric matrix routes through eigh and PROMOTES the real spectrum513
// and eigenvectors to complex (imag 0): values are real-valued complex, and the514
// eigenvectors are present (a 2x2 complex matrix), unlike the non-symmetric case.515
const la::EigC e = la::eig(mat(2, 2, {2, 1, 1, 2})); // eigenvalues 3, 1516
EXPECT_TRUE(cclose(cget(e.values, {0}), 3.0));517
EXPECT_TRUE(cclose(cget(e.values, {1}), 1.0));518
EXPECT_EQ(nd::size_of(e.vectors), 4); // 2x2 eigenvectors present (promoted to complex)519
EXPECT_EQ(e.vectors.ndim(), 2u);520
}522
// ---- targeted tests for the deep numerical branches ----524
namespace {525
std::vector<double> sorted3(const nd::NDArray& v) {526
std::vector<double> s{nd::get(v, {0}), nd::get(v, {1}), nd::get(v, {2})};527
std::sort(s.begin(), s.end());528
return s;529
}530
// Same, for a complex spectrum known to be real (imaginary parts ≈ 0): the real parts.531
std::vector<double> sorted3c(const la::CNDArray& v) {532
std::vector<double> s{cget(v, {0}).real(), cget(v, {1}).real(), cget(v, {2}).real()};533
std::sort(s.begin(), s.end());534
return s;535
}536
} // namespace538
TEST(LinalgRoutines, VdotInnerAcceptTwoDimVectors) {539
// A 2-D Nx1 / 1xN is treated as a flat vector by vdot/inner.540
EXPECT_DOUBLE_EQ(la::vdot(mat(3, 1, {1, 2, 3}), mat(3, 1, {4, 5, 6})), 32.0);541
EXPECT_DOUBLE_EQ(la::inner(mat(1, 3, {1, 2, 3}), mat(1, 3, {4, 5, 6})), 32.0);542
}544
TEST(LinalgRoutines, DetRequiresPivot) {545
EXPECT_TRUE(close(la::det(mat(2, 2, {0, 1, 1, 0})), -1.0)); // forces an LU row swap546
}548
TEST(LinalgRoutines, Eigvalsh3x3Dense) {549
// [[2,1,1],[1,2,1],[1,1,2]] = I + ones -> eigenvalues 4, 1, 1 (Jacobi rotations).550
const std::vector<double> v = sorted3(la::eigvalsh(mat(3, 3, {2, 1, 1, 1, 2, 1, 1, 1, 2})));551
EXPECT_TRUE(close(v[0], 1.0, 1e-6));552
EXPECT_TRUE(close(v[1], 1.0, 1e-6));553
EXPECT_TRUE(close(v[2], 4.0, 1e-6));554
}556
TEST(LinalgRoutines, GeneralEigvals3x3Dense) {557
// Same dense matrix through the general (Hessenberg + shifted-QR) path.558
const std::vector<double> v = sorted3c(la::eigvals(mat(3, 3, {2, 1, 1, 1, 2, 1, 1, 1, 2})));559
EXPECT_TRUE(close(v[0], 1.0, 1e-6));560
EXPECT_TRUE(close(v[2], 4.0, 1e-6));561
}563
TEST(LinalgRoutines, NonSymmetric3x3HessenbergPath) {564
// M = P·diag(2,3,5)·P⁻¹ — non-symmetric with real eigenvalues 2,3,5. Routes565
// through eigvals_general (Householder–Hessenberg + shifted QR for n≥3).566
const nd::NDArray M = mat(3, 3, {2.5, 0.5, -0.5, -1, 4, 1, -1.5, 1.5, 3.5});567
const std::vector<double> v = sorted3c(la::eigvals(M));568
EXPECT_TRUE(close(v[0], 2.0, 1e-6));569
EXPECT_TRUE(close(v[1], 3.0, 1e-6));570
EXPECT_TRUE(close(v[2], 5.0, 1e-6));571
EXPECT_TRUE(cclose(cget(la::eig(M).values, {0}), 5.0)); // descending; eig() too572
}574
TEST(LinalgRoutines, ComplexEigenvaluesOfRotation) {575
// A 2-D rotation [[0,-1],[1,0]] has eigenvalues ±i. The general eigensolver576
// returns the complex conjugate pair (descending by real, then imag: +i, then -i)577
// rather than throwing — complex spectra are first-class.578
const la::CNDArray ev = la::eigvals(mat(2, 2, {0, -1, 1, 0}));579
EXPECT_TRUE(cclose(cget(ev, {0}), std::complex<double>(0.0, 1.0)));580
EXPECT_TRUE(cclose(cget(ev, {1}), std::complex<double>(0.0, -1.0)));581
// A complex pair with a non-zero real part: [[1,-1],[1,1]] -> 1±i.582
const la::CNDArray ev2 = la::eigvals(mat(2, 2, {1, -1, 1, 1}));583
EXPECT_TRUE(cclose(cget(ev2, {0}), std::complex<double>(1.0, 1.0)));584
EXPECT_TRUE(cclose(cget(ev2, {1}), std::complex<double>(1.0, -1.0)));585
}587
TEST(LinalgRoutines, VdotRejectsNonVector) {588
EXPECT_THROW(static_cast<void>(la::vdot(mat(2, 2, {1, 2, 3, 4}), mat(2, 2, {1, 2, 3, 4}))), std::runtime_error);589
}591
TEST(LinalgRoutines, NormOfMatrixIsFrobenius) {592
EXPECT_TRUE(close(la::norm(mat(2, 2, {1, 2, 2, 4})), 5.0)); // sqrt(1+4+4+16)593
}595
TEST(LinalgRoutines, PinvCondRankOnWideMatrix) {596
const nd::NDArray W = mat(2, 3, {1, 0, 0, 0, 1, 0}); // 2x3 (more cols than rows)597
const nd::NDArray P = la::pinv(W); // -> 3x2 (transpose-SVD branch)598
EXPECT_EQ(nd::shape_of(P), (std::vector<long long>{3, 2}));599
EXPECT_GE(la::cond(W), 1.0);600
EXPECT_EQ(la::matrix_rank(W), 2);601
}603
// ---- coverage: non-contiguous inputs, edge branches, and defensive throws ----604
TEST(LinalgRoutines, NonContiguousInputs) {605
// A broadcast (stride-0) view is non-contiguous, exercising the packing fallback in606
// contig/as_matrix/as_vector/as_cmatrix (the contiguous fast path runs everywhere else).607
const nd::NDArray ncvec = nd::broadcast_to(nd::scalar(2.0), {4}); // [2,2,2,2]608
EXPECT_TRUE(close(la::dot(ncvec, ncvec), 16.0, 1e-12)); // contig() pack path609
EXPECT_TRUE(close(la::norm(nd::broadcast_to(nd::scalar(3.0), {2, 2})), 6.0, 1e-12));610
const nd::NDArray ncmat = nd::broadcast_to(nd::array({1.0, 2.0, 3.0}), {3, 3});611
EXPECT_TRUE(close(la::det(ncmat), 0.0, 1e-9)); // as_matrix pack path612
const nd::NDArray I3 = mat(3, 3, {1, 0, 0, 0, 1, 0, 0, 0, 1});613
const nd::NDArray x = la::solve(I3, nd::broadcast_to(nd::scalar(5.0), {3})); // as_vector pack614
EXPECT_TRUE(close(nd::get(x, {0}), 5.0, 1e-9));615
const la::CNDArray nccx = nd::broadcast_to(nd::scalar(C(1, 1)), {2, 2});616
EXPECT_EQ(la::conj_transpose(nccx).ndim(), 2u); // complex contig() pack617
// complex eigvalsh extracts via as_cmatrix — a broadcast [[2,2],[2,2]] (Hermitian) packs.618
EXPECT_TRUE(close(nd::get(la::eigvalsh(nd::broadcast_to(nd::scalar(C(2, 0)), {2, 2})), {0}),619
4.0, 1e-9));620
}622
TEST(LinalgRoutines, ComplexDotFourPlusElements) {623
// 4+ elements drives the multi-accumulator loop in cdot (both dot and the conjugating vdot).624
const la::CNDArray a = cvec({C(1, 1), C(2, 0), C(0, 1), C(1, -1), C(2, 2)});625
const la::CNDArray b = cvec({C(1, 0), C(0, 1), C(1, 1), C(2, 0), C(0, -1)});626
C dotref{}, vdotref{};627
for (long long i = 0; i < 5; ++i) {628
const C ai = cget(a, {i}), bi = cget(b, {i});629
dotref += ai * bi;630
vdotref += std::conj(ai) * bi;631
}632
EXPECT_TRUE(cclose(la::dot(a, b), dotref)); // non-conjugating multi-accumulator633
EXPECT_TRUE(cclose(la::vdot(a, b), vdotref)); // conjugating multi-accumulator634
}636
TEST(LinalgRoutines, ShapeAndConvergenceGuards) {637
const nd::NDArray v = nd::array({1.0, 2.0});638
EXPECT_THROW(static_cast<void>(la::matmul(v, v)), std::runtime_error); // real matmul non-2D639
EXPECT_THROW(static_cast<void>(la::matmul(mat(2, 3, {1, 2, 3, 4, 5, 6}), mat(2, 2, {1, 2, 3, 4}))),640
std::runtime_error); // matmul front inner-dim mismatch (3 != 2)641
EXPECT_THROW(static_cast<void>(la::kron(v, v)), std::runtime_error); // kron non-2D642
const la::CNDArray cv = cvec({C(1, 0), C(2, 0)});643
EXPECT_THROW(static_cast<void>(la::matmul(cv, cv)), std::runtime_error); // complex matmul non-2D644
// inv that requires a row pivot (zero leading pivot)645
const nd::NDArray inv = la::inv(mat(2, 2, {0, 1, 1, 0}));646
EXPECT_TRUE(close(nd::get(inv, {0, 1}), 1.0, 1e-9));647
// values-only SVD on a WIDE matrix routes through the transpose branch648
EXPECT_TRUE(close(nd::get(la::svdvals(mat(2, 3, {1, 0, 0, 0, 1, 0})), {0}), 1.0, 1e-9));649
// NaN input never converges -> the defensive "did not converge" throws fire650
const double nan = std::numeric_limits<double>::quiet_NaN();651
EXPECT_THROW(static_cast<void>(la::eigvalsh(mat(2, 2, {nan, 0, 0, 1}))), std::runtime_error); // symmetric QL652
EXPECT_THROW(static_cast<void>(la::svdvals(mat(2, 2, {nan, 0, 0, 1}))), std::runtime_error); // SVD QR653
}655
TEST(LinalgRoutines, RankDeficientAndDiagonalPaths) {656
// A matrix with an exactly-zero column gives an exactly-zero singular value, which657
// drives the g==0 branch in U accumulation and the bulge cancellation in the SVD QR.658
const la::SVD s = la::svd(mat(3, 3, {1, 4, 0, 2, 5, 0, 3, 6, 0}));659
EXPECT_TRUE(close(nd::get(s.s, {2}), 0.0, 1e-9));660
// A diagonal symmetric matrix has all-zero off-diagonals -> the scale==0 branch in tred2.661
const la::Eig e = la::eigh(mat(3, 3, {2, 0, 0, 0, 5, 0, 0, 0, 7}));662
EXPECT_TRUE(close(nd::get(e.values, {0}), 7.0, 1e-9));663
}665
TEST(LinalgRoutines, SvdCancellationPath) {666
// An upper-bidiagonal matrix with an interior zero diagonal (w[1]=0) but a667
// non-negligible super-diagonal triggers the QR "cancel rv1" Givens sweep in U.668
(void)la::svd(mat(3, 3, {2, 3, 0, 0, 0, 4, 0, 0, 5}));669
(void)la::svd(mat(3, 3, {0, 5, 0, 0, 0, 5, 0, 0, 0}));670
(void)la::svd(mat(4, 4, {1, 9, 0, 0, 0, 0, 9, 0, 0, 0, 0, 9, 0, 0, 0, 1}));671
SUCCEED();672
}674
// Exercise the WIDE-UNROLL main loops of the multi-accumulator/blocked kernels: the675
// existing tests use tiny matrices that only ever run the scalar remainder, leaving the676
// 4-/8-wide vectorized bodies (ddot, real+complex matmul row-blocking, cholesky/qr/677
// tred2/trace reductions) uncovered. These use n≥8 so the main loops run.678
TEST(LinalgRoutines, WideKernelPaths) {679
// 8×8 SPD: diag 10, off-diag 1 (= 9·I + J). Eigenvalues {17, 9×7}, trace 80.680
std::vector<double> a8(64);681
for (std::size_t i = 0; i < 8; ++i)682
for (std::size_t j = 0; j < 8; ++j) a8[i * 8 + j] = (i == j) ? 10.0 : 1.0;683
const nd::NDArray A = mat(8, 8, a8);685
// dot over ≥8 elements → ddot 8-wide body.686
EXPECT_DOUBLE_EQ(la::dot(nd::array(std::vector<double>(10, 1.0)),687
nd::array(std::vector<double>(10, 2.0))), 20.0);688
// trace 8×8 → trace 4-wide body.689
EXPECT_DOUBLE_EQ(la::trace(A), 80.0);690
// matmul 8×8 → real 4-row block. A·I == A.691
std::vector<double> id8(64, 0.0);692
for (std::size_t i = 0; i < 8; ++i) id8[i * 8 + i] = 1.0;693
const nd::NDArray AI = la::matmul(A, mat(8, 8, id8));694
EXPECT_DOUBLE_EQ(nd::get(AI, {0, 0}), 10.0);695
EXPECT_DOUBLE_EQ(nd::get(AI, {1, 0}), 1.0);696
// cholesky 8×8 (j reaches ≥4 → 4-wide inner dot). Reconstruct A = L·Lᵀ.697
const nd::NDArray L = la::cholesky(A);698
double a00 = 0;699
for (long long k = 0; k < 8; ++k) a00 += nd::get(L, {0, k}) * nd::get(L, {0, k});700
EXPECT_NEAR(a00, 10.0, 1e-9);701
// qr 8×4 → reflect 4-wide body. R upper-triangular, Q·R == A_panel.702
std::vector<double> p(32);703
for (std::size_t i = 0; i < 8; ++i)704
for (std::size_t j = 0; j < 4; ++j) p[i * 4 + j] = a8[i * 8 + j];705
const la::QR qr = la::qr(mat(8, 4, p));706
EXPECT_NEAR(nd::get(qr.r, {1, 0}), 0.0, 1e-9); // upper-triangular707
// eigvalsh 8×8 → tred2 mat-vec 4-wide body. Largest eigenvalue 17, sum 80.708
const nd::NDArray w = la::eigvalsh(A);709
EXPECT_NEAR(nd::get(w, {0}), 17.0, 1e-7);710
double sw = 0;711
for (long long i = 0; i < 8; ++i) sw += nd::get(w, {i});712
EXPECT_NEAR(sw, 80.0, 1e-7);713
// eig() on a symmetric matrix → the reuse-the-extracted-A symmetric branch.714
const la::EigC e = la::eig(A);715
EXPECT_NEAR(cget(e.values, {0}).real(), 17.0, 1e-6);717
// complex matmul 8×8 → complex 4-row block.718
std::vector<C> z(64), zi(64, C{0, 0});719
for (std::size_t i = 0; i < 8; ++i) { z[i * 8 + i] = C{2, 0}; zi[i * 8 + i] = C{1, 0}; }720
const la::CNDArray Z = cmat(8, 8, z), I = cmat(8, 8, zi);721
EXPECT_TRUE(cclose(cget(la::matmul(Z, I), {3, 3}), C{2, 0}));722
}724
// Non-contiguous (broadcast/strided) operands take the scratch-packing fallback in the725
// products/reductions, not the zero-copy fast path.726
TEST(LinalgRoutines, StridedOperandFallback) {727
const nd::NDArray s = nd::broadcast_to(nd::scalar(2.0), {10}); // stride-0 view, len 10728
EXPECT_DOUBLE_EQ(la::dot(s, s), 40.0); // 10 · (2·2)729
EXPECT_NEAR(la::norm(s), std::sqrt(40.0), 1e-9);730
const la::CNDArray cs = nd::broadcast_to(nd::scalar(C{2, 0}), {6});731
EXPECT_TRUE(cclose(la::dot(cs, cs), C{24, 0})); // 6·(2·2)732
EXPECT_TRUE(cclose(la::vdot(cs, cs), C{24, 0})); // conj path, strided733
}735
// Batched matmul: [B,M,K] @ [B,K,N] -> [B,M,N], each slice the same product the 2-D kernel736
// gives; strict batching (mismatched batch counts / mixed ranks throw).737
TEST(LinalgRoutines, BatchedMatmul) {738
// Two batches of 2x3 @ 3x2, second batch = 2x the first: slices must match the 2-D results.739
std::vector<double> av{1, 2, 3, 4, 5, 6}, bv{7, 8, 9, 10, 11, 12};740
std::vector<double> abatch(av), bbatch(bv);741
for (double x : av) abatch.push_back(2 * x); // batch 1 doubles A742
for (double x : bv) bbatch.push_back(x); // batch 1 reuses B743
const nd::NDArray A = nd::reshape(nd::array(std::move(abatch)), {2, 2, 3});744
const nd::NDArray B = nd::reshape(nd::array(std::move(bbatch)), {2, 3, 2});745
const nd::NDArray C3 = la::matmul(A, B);746
ASSERT_EQ(C3.ndim(), 3u);747
EXPECT_EQ(C3.shape()[0], 2u);748
EXPECT_EQ(C3.shape()[1], 2u);749
EXPECT_EQ(C3.shape()[2], 2u);750
// slice 0: the classic {58 64; 139 154}; slice 1 doubles it.751
EXPECT_DOUBLE_EQ(nd::get(C3, {0, 0, 0}), 58.0);752
EXPECT_DOUBLE_EQ(nd::get(C3, {0, 1, 1}), 154.0);753
EXPECT_DOUBLE_EQ(nd::get(C3, {1, 0, 0}), 116.0);754
EXPECT_DOUBLE_EQ(nd::get(C3, {1, 1, 1}), 308.0);756
// The out-param kernel form reuses the caller's 3-D buffer.757
nd::NDArray out = nd::zeros({2, 2, 2});758
la::matmul(out, A, B);759
EXPECT_DOUBLE_EQ(nd::get(out, {1, 0, 1}), 2 * 64.0);761
// Strictness: mixed rank and batch-count mismatch throw.762
EXPECT_THROW(static_cast<void>(la::matmul(A, nd::zeros({3, 2}))), std::runtime_error);763
EXPECT_THROW(static_cast<void>(la::matmul(A, nd::zeros({3, 3, 2}))), std::runtime_error);764
}766
// The REAL instantiation of the (conjugate-)transpose kernel. The complex path is covered by767
// ComplexProducts, but a real element takes the other side of the kernel's `if constexpr` — a plain768
// transpose with the conjugation compiled out — and nothing exercised it, so the branch that most769
// users actually hit was the untested one.770
TEST(LinalgRoutines, ConjTransposeOnRealMatrix) {771
const nd::NDArray M = mat(2, 3, {1, 2, 3, 4, 5, 6});772
const nd::NDArray T = la::conj_transpose(M);773
ASSERT_EQ(T.ndim(), 2);774
EXPECT_EQ(T.shape()[0], 3u);775
EXPECT_EQ(T.shape()[1], 2u);776
EXPECT_TRUE(close(nd::get(T, {0, 0}), 1));777
EXPECT_TRUE(close(nd::get(T, {0, 1}), 4));778
EXPECT_TRUE(close(nd::get(T, {2, 0}), 3));779
EXPECT_TRUE(close(nd::get(T, {2, 1}), 6));781
// Same answer through the allocation-free out-param form.782
nd::NDArray out = nd::NDArray::uninitialized({3, 2});783
la::conj_transpose(out, M);784
EXPECT_TRUE(close(nd::get(out, {0, 1}), 4));785
EXPECT_TRUE(close(nd::get(out, {2, 0}), 3));786
}788
// The out-param shape guard on the COPY path. `outer` validates its out against an initializer-list789
// shape; the routines that build a result and then copy it in (cholesky, inv, solve, lstsq) validate790
// against a runtime shape vector instead — a separate overload, and one no test reached. A wrong-shaped791
// out must be refused rather than write past the caller's buffer.792
TEST(LinalgRoutines, OutParamRejectsWrongShapeOnTheCopyPath) {793
const nd::NDArray spd = mat(2, 2, {4, 2, 2, 3}); // symmetric positive-definite794
nd::NDArray good = nd::NDArray::uninitialized({2, 2});795
EXPECT_NO_THROW(la::cholesky(good, spd));797
nd::NDArray wrong = nd::NDArray::uninitialized({3, 3});798
EXPECT_THROW(static_cast<void>(la::cholesky(wrong, spd)), std::runtime_error);799
}