Files
tensor/linalg/linalg_complex_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

262 lines
6.8 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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)
}
}