feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,261 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package linalg
|
||||
|
||||
import (
|
||||
"math"
|
||||
"math/cmplx"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// complexMatrix builds a complex128 array from row-major values.
|
||||
func complexMatrix(t *testing.T, shape []int, vals ...complex128) *core.Array {
|
||||
t.Helper()
|
||||
total := 1
|
||||
for _, d := range shape {
|
||||
total *= d
|
||||
}
|
||||
if len(vals) != total {
|
||||
t.Fatalf("value count %d does not fill %v", len(vals), shape)
|
||||
}
|
||||
a, err := core.ComplexFromArray(vals, shape...)
|
||||
if err != nil {
|
||||
t.Fatalf("ComplexFromArray(%v, %v): %v", vals, shape, err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
func TestDetComplexRotation(t *testing.T) {
|
||||
// A rotation in the complex plane has unit determinant; scaling one
|
||||
// row by (2+i) multiplies the determinant by the same factor.
|
||||
angle := complex(0.6, -0.8) // unit magnitude
|
||||
m := complexMatrix(t, []int{2, 2},
|
||||
angle, 0,
|
||||
0, 1,
|
||||
)
|
||||
det, err := DetComplex(m)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// The complex diagonal keeps its phase, so the unit-modulus claim
|
||||
// holds on the magnitude only.
|
||||
if math.Abs(cmplx.Abs(det)-1) > 1e-12 {
|
||||
t.Fatalf("unitary |det| = %v", cmplx.Abs(det))
|
||||
}
|
||||
|
||||
scaled := complexMatrix(t, []int{2, 2},
|
||||
2+1i, 0,
|
||||
3-4i, angle,
|
||||
)
|
||||
det2, err := DetComplex(scaled)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := (2 + 1i) * angle
|
||||
if cmplx.Abs(det2-want) > 1e-12 {
|
||||
t.Fatalf("scaled det = %v, want %v", det2, want)
|
||||
}
|
||||
|
||||
// A singular matrix yields zero without erroring.
|
||||
singular := complexMatrix(t, []int{2, 2}, 1, 2i, 2, 4i)
|
||||
detS, err := DetComplex(singular)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if detS != 0 {
|
||||
t.Fatalf("singular det = %v, want 0", detS)
|
||||
}
|
||||
|
||||
// The real Det keeps pointing complex callers at the complex twin.
|
||||
realShaped, _ := core.FromFloats([]float64{1, 0, 0, 1}, 2, 2)
|
||||
if _, err := Det(realShaped); err != nil {
|
||||
t.Fatalf("real det errored: %v", err)
|
||||
}
|
||||
if _, err := Det(complexMatrix(t, []int{2, 2}, 1, 0, 0, 1)); err == nil {
|
||||
t.Fatal("real Det accepted a complex matrix")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSolveInvComplexRoundTrip(t *testing.T) {
|
||||
a := complexMatrix(t, []int{3, 3},
|
||||
2+1i, 0, 1-1i,
|
||||
0, 1-2i, 0.5,
|
||||
1i, 1, 3+2i,
|
||||
)
|
||||
|
||||
inv, err := Inv(a)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if inv.Dtype() != core.Complex {
|
||||
t.Fatalf("inverse dtype %v", inv.Dtype())
|
||||
}
|
||||
product, err := core.MatMul2D(a, inv)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := range 9 {
|
||||
row, col := i/3, i%3
|
||||
want := complex(0, 0)
|
||||
if row == col {
|
||||
want = 1
|
||||
}
|
||||
if got := product.RawComplexes()[i]; cmplx.Abs(got-want) > 1e-10 {
|
||||
t.Fatalf("A·A⁻¹[%d][%d] = %v, want %v", row, col, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// Solve and verify by substitution for a vector right-hand side.
|
||||
b, _ := core.FromFloats([]float64{1, 2, 3}, 3) // real b promotes to complex
|
||||
x, err := Solve(a, b)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
check, err := core.MatMul2D(a, x)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := range 3 {
|
||||
if got := check.RawComplexes()[i]; cmplx.Abs(got-complex(float64(i+1), 0)) > 1e-10 {
|
||||
t.Fatalf("(a·x)[%d] = %v, want %v", i, got, float64(i+1))
|
||||
}
|
||||
}
|
||||
|
||||
// Solving with an explicitly complex vector works symmetrically.
|
||||
bc := complexMatrix(t, []int{3}, 2+2i, 0, -1i)
|
||||
bcol, err := core.Reshape(bc, 3)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
xc, err := Solve(a, bcol)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if xc.Dtype() != core.Complex || xc.Len() != 3 {
|
||||
t.Fatalf("complex solve shape/dtype: %v %v", xc.Shape(), xc.Dtype())
|
||||
}
|
||||
|
||||
// Singular systems error instead of returning garbage.
|
||||
singular := complexMatrix(t, []int{2, 2}, 1, 1i, 2, 2i)
|
||||
bad, _ := core.FromFloats([]float64{1, 1}, 2)
|
||||
if _, err := Solve(singular, bad); err == nil {
|
||||
t.Fatal("singular solve succeeded")
|
||||
}
|
||||
if _, err := Inv(singular); err == nil {
|
||||
t.Fatal("singular inverse succeeded")
|
||||
}
|
||||
}
|
||||
|
||||
func TestKronComplexBlocks(t *testing.T) {
|
||||
a := complexMatrix(t, []int{2, 2}, 1+1i, 0, 0, 2-1i)
|
||||
b := complexMatrix(t, []int{2, 2}, 1, 2i, 3, 0)
|
||||
|
||||
out, err := core.Kron(a, b)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if out.Dtype() != core.Complex {
|
||||
t.Fatalf("kron dtype %v", out.Dtype())
|
||||
}
|
||||
if got := out.Shape(); got[0] != 4 || got[1] != 4 {
|
||||
t.Fatalf("kron shape %v", got)
|
||||
}
|
||||
// Top-left block scales b by (1+1i); it sits on flat positions
|
||||
// k*4+l because blocks interleave in the outer product layout.
|
||||
type slot struct {
|
||||
idx int
|
||||
want complex128
|
||||
}
|
||||
for _, s := range []slot{
|
||||
{0, 1 + 1i},
|
||||
{1, 2i * (1 + 1i)},
|
||||
{4, 3 * (1 + 1i)},
|
||||
{5, 0},
|
||||
} {
|
||||
if got := out.RawComplexes()[s.idx]; got != s.want {
|
||||
t.Fatalf("block TL[%d] = %v, want %v", s.idx, got, s.want)
|
||||
}
|
||||
}
|
||||
// Bottom-right block scales b by (2−1i); identity via mixed pair.
|
||||
mixed, _ := core.FromFloats([]float64{1}, 1, 1) // real 1×1
|
||||
eye := complexMatrix(t, []int{1, 1}, 2-1i)
|
||||
cross, err := core.Kron(mixed, eye)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if cross.RawComplexes()[0] != 2-1i {
|
||||
t.Fatalf("mixed kron element = %v", cross.RawComplexes()[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestTraceComplexSum(t *testing.T) {
|
||||
m := complexMatrix(t, []int{2, 2}, 1+2i, 99, 99, 3-4i)
|
||||
s, err := core.TraceComplex(m)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if s != 4-2i {
|
||||
t.Fatalf("trace = %v, want 4−2i", s)
|
||||
}
|
||||
nonSquare := complexMatrix(t, []int{1, 2}, 1, 2)
|
||||
if _, err := core.TraceComplex(nonSquare); err == nil {
|
||||
t.Fatal("non-square trace accepted")
|
||||
}
|
||||
// Real path keeps rejecting complexes by name.
|
||||
realM, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
|
||||
if _, err := core.TraceComplex(realM); err == nil {
|
||||
t.Fatal("TraceComplex accepted a real matrix")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPowIComplexExactness(t *testing.T) {
|
||||
z := mustComplexes(t, []complex128{1 + 1i, 2 - 1i}, 2)
|
||||
|
||||
cubed, err := core.PowI(z, 3)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
wantFirst := (1 + 1i) * (1 + 1i) * (1 + 1i) // −2+2i
|
||||
if cubed.RawComplexes()[0] != wantFirst {
|
||||
t.Fatalf("(1+i)^3 = %v, want %v", cubed.RawComplexes()[0], wantFirst)
|
||||
}
|
||||
inverse, err := core.PowI(z, -1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
one, err := core.Mul(inverse, z)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := range 2 {
|
||||
if math.Abs(real(one.RawComplexes()[i])-1) > 1e-12 || math.Abs(imag(one.RawComplexes()[i])) > 1e-12 {
|
||||
t.Fatalf("z·z⁻¹[%d] = %v", i, one.RawComplexes()[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestComplexSolveInvDetErrors moved with Solve, Inv and Det from the
|
||||
// root package: Solve and Inv run on complex systems through the same
|
||||
// LU kernel, while the real-only Det points at its complex twin.
|
||||
func TestComplexSolveInvDetErrors(t *testing.T) {
|
||||
sys := complexMatrix(t, []int{2, 2},
|
||||
2+1i, 0,
|
||||
0, 3-1i,
|
||||
)
|
||||
got, err := Inv(sys)
|
||||
if err != nil {
|
||||
t.Fatalf("Inv complex: %v", err)
|
||||
}
|
||||
if got.Dtype() != core.Complex {
|
||||
t.Fatalf("Inv complex dtype %v", got.Dtype())
|
||||
}
|
||||
if _, err := Solve(sys, sys); err != nil {
|
||||
t.Fatalf("Solve complex: %v", err)
|
||||
}
|
||||
if _, err := Det(sys); err == nil || !strings.Contains(err.Error(), "DetComplex") {
|
||||
t.Fatalf("Det complex: %v", err)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user