Files
tensor/linalg/complex_refusal_pins_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

171 lines
5.0 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"
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Regression tests: entry points that reached a
// complex array's float accessor (a panic, not an error), a
// preconditioner applied to a system of another dimension, and Krylov
// callbacks returning the wrong length.
// mustComplex builds a complex array.
func mustComplex(t *testing.T, vals []complex128, shape ...int) *core.Array {
t.Helper()
a, err := core.FromComplexes(vals, shape...)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
return a
}
// TestComplexInputsAreRefused pins the dtype gate of the four entry
// points that read their arguments as real numbers.
func TestComplexInputsAreRefused(t *testing.T) {
c2 := mustComplex(t, []complex128{1 + 1i, 1 + 1i}, 2)
c3 := mustComplex(t, []complex128{1 + 1i, 2 + 2i, 3 + 3i}, 3)
t.Run("SolveTridiagonal", func(t *testing.T) {
if _, err := SolveTridiagonal(c2, c3, c2, c3); err == nil {
t.Fatal("expected an error for complex diagonals")
} else if !strings.Contains(err.Error(), "complex") {
t.Fatalf("error = %v, want a complex-dtype refusal", err)
}
})
t.Run("SolveCyclicTridiagonal", func(t *testing.T) {
if _, err := SolveCyclicTridiagonal(c3, c3, c3, c3); err == nil {
t.Fatal("expected an error for complex diagonals")
}
})
t.Run("FitPolynomial", func(t *testing.T) {
if _, err := FitPolynomial(c2, c2, 1); err == nil {
t.Fatal("expected an error for complex samples")
}
})
t.Run("NewCubicSpline", func(t *testing.T) {
if _, err := NewCubicSpline(c3, c3); err == nil {
t.Fatal("expected an error for complex knots")
}
})
}
// TestSparseILUDimensionMismatch pins the preconditioner check: an ILU
// built for another system must be refused rather than indexed.
func TestSparseILUDimensionMismatch(t *testing.T) {
mk := func(n int, vals []float64) *core.SparseCOO {
t.Helper()
idx := make([]int64, 0, n*2)
for i := range n {
idx = append(idx, int64(i), int64(i))
}
ia, err := core.FromInts(idx, n, 2)
if err != nil {
t.Fatal(err)
}
va, err := core.FromFloats(vals, n)
if err != nil {
t.Fatal(err)
}
sp, err := core.NewSparseCOO(ia, va, []int{n, n})
if err != nil {
t.Fatal(err)
}
return sp
}
big := mk(4, []float64{4, 4, 4, 4})
small := mk(3, []float64{3, 3, 3})
ilu, err := NewSparseILU(small)
if err != nil {
t.Fatalf("NewSparseILU: %v", err)
}
b := mustFloats(t, []float64{1, 1, 1, 1}, 4)
if _, err := SpSolve(big, b, 0, 50, ilu); err == nil {
t.Fatal("expected an error for a preconditioner of the wrong dimension")
} else if !strings.Contains(err.Error(), "dimension") {
t.Fatalf("error = %v, want a dimension refusal", err)
}
if _, err := SpSolveBiCGSTAB(big, b, 0, 50, ilu); err == nil {
t.Fatal("expected an error from BiCGSTAB for a preconditioner of the wrong dimension")
}
}
// TestGMRESCallbackLength pins the op contract: a callback that
// returns fewer elements than the problem has must be an error.
func TestGMRESCallbackLength(t *testing.T) {
b := mustFloats(t, []float64{1, 2}, 2)
short := mustFloats(t, []float64{1}, 1)
if _, err := GMRES(func(*core.Array) (*core.Array, error) { return short, nil }, b, 2, 5, 1e-12); err == nil {
t.Fatal("expected an error for an op returning the wrong length")
} else if !strings.Contains(err.Error(), "elements") {
t.Fatalf("error = %v, want a length refusal", err)
}
}
// TestQRLargeScale pins the Householder reflector at scales where the
// raw squares of the old norm overflowed: every intermediate of the
// reflector is now O(1) in scaled units, so the reconstruction holds at
// any magnitude.
func TestQRLargeScale(t *testing.T) {
for _, scale := range []float64{1, 1e10, 1e100, 1.2589e154, 1e200, 1e300} {
vals := []float64{2, -1, 0.5, 1, 3, -2, 4, 1, 1}
for i := range vals {
vals[i] *= scale
}
a, err := core.FromFloats(vals, 3, 3)
if err != nil {
t.Fatal(err)
}
q, r, err := QR(a)
if err != nil {
t.Fatalf("scale %g: %v", scale, err)
}
maxA := 0.0
for _, v := range vals {
maxA = math.Max(maxA, math.Abs(v))
}
worst := 0.0
for i := range 3 {
for j := range 3 {
acc := 0.0
for k := range 3 {
acc += q.FloatAt(i*3+k) * r.FloatAt(k*3+j)
}
worst = math.Max(worst, math.Abs(acc-vals[i*3+j])/maxA)
}
}
if !(worst < 1e-12) {
t.Errorf("scale %g: |QR−A|/|A| = %g, want below 1e-12", scale, worst)
}
}
}
// TestSparseILUNil pins the nil guard: an explicitly nil preconditioner
// is an error, not a dereference.
func TestSparseILUNil(t *testing.T) {
a, err := core.FromFloats([]float64{2, 3}, 2)
if err != nil {
t.Fatal(err)
}
idx, err := core.FromInts([]int64{0, 0, 1, 1}, 2, 2)
if err != nil {
t.Fatal(err)
}
sp, err := core.NewSparseCOO(idx, a, []int{2, 2})
if err != nil {
t.Fatal(err)
}
b, err := core.FromFloats([]float64{1, 1}, 2)
if err != nil {
t.Fatal(err)
}
if _, err := SpSolve(sp, b, 0, 10, nil); err == nil {
t.Fatal("expected an error for a nil preconditioner")
}
}