feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,167 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package linalg
|
||||
|
||||
import (
|
||||
"math"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGMRES(t *testing.T) {
|
||||
vals := []float64{
|
||||
10, 1, 0, 2,
|
||||
1, 12, 3, 0,
|
||||
0, 3, 15, 1,
|
||||
2, 0, 1, 8,
|
||||
}
|
||||
op := func(v *core.Array) (*core.Array, error) {
|
||||
out := core.New(core.Float, 4)
|
||||
for i := range 4 {
|
||||
s := 0.0
|
||||
for j := range 4 {
|
||||
s += vals[i*4+j] * v.FloatAt(j)
|
||||
}
|
||||
out.RawFloats()[i] = s
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
b := mustFloats(t, []float64{1, 2, 3, 4}, 4)
|
||||
x, err := GMRES(op, b, 0, 0, 1e-12)
|
||||
if err != nil {
|
||||
t.Fatalf("GMRES: %v", err)
|
||||
}
|
||||
a := mustFloats(t, vals, 4, 4)
|
||||
ref, _ := Solve(a, b)
|
||||
for i := range 4 {
|
||||
if math.Abs(x.FloatAt(i)-ref.FloatAt(i)) > 1e-10 {
|
||||
t.Fatalf("x[%d] = %.16g, want %.16g", i, x.FloatAt(i), ref.FloatAt(i))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGMRESLarger(t *testing.T) {
|
||||
const n = 50
|
||||
vals := make([]float64, n*n)
|
||||
for i := range n {
|
||||
vals[i*n+i] = 4
|
||||
if i+1 < n {
|
||||
vals[i*n+i+1] = -1
|
||||
vals[(i+1)*n+i] = -1
|
||||
}
|
||||
}
|
||||
op := func(v *core.Array) (*core.Array, error) {
|
||||
out := core.New(core.Float, n)
|
||||
for i := range n {
|
||||
s := 0.0
|
||||
for j := range n {
|
||||
s += vals[i*n+j] * v.FloatAt(j)
|
||||
}
|
||||
out.RawFloats()[i] = s
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
bv := make([]float64, n)
|
||||
for i := range n {
|
||||
bv[i] = float64(i + 1)
|
||||
}
|
||||
b := mustFloats(t, bv, n)
|
||||
x, err := GMRES(op, b, 0, 0, 1e-12)
|
||||
if err != nil {
|
||||
t.Fatalf("GMRES: %v", err)
|
||||
}
|
||||
for i := range n {
|
||||
s := 0.0
|
||||
for j := range n {
|
||||
s += vals[i*n+j] * x.FloatAt(j)
|
||||
}
|
||||
if math.Abs(s-bv[i]) > 1e-8 {
|
||||
t.Fatalf("res[%d] = %e, want < 1e-8", i, math.Abs(s-bv[i]))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestGMRESRestarted forces several outer cycles by capping the
|
||||
// Krylov dimension below the system size; the accumulated solution
|
||||
// must still match the dense solve.
|
||||
func TestGMRESRestarted(t *testing.T) {
|
||||
vals := []float64{
|
||||
10, 1, 0, 2,
|
||||
1, 12, 3, 0,
|
||||
0, 3, 15, 1,
|
||||
2, 0, 1, 8,
|
||||
}
|
||||
op := func(v *core.Array) (*core.Array, error) {
|
||||
out := core.New(core.Float, 4)
|
||||
for i := range 4 {
|
||||
s := 0.0
|
||||
for j := range 4 {
|
||||
s += vals[i*4+j] * v.FloatAt(j)
|
||||
}
|
||||
out.RawFloats()[i] = s
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
b := mustFloats(t, []float64{1, 2, 3, 4}, 4)
|
||||
a := mustFloats(t, vals, 4, 4)
|
||||
ref, err := Solve(a, b)
|
||||
if err != nil {
|
||||
t.Fatalf("Solve: %v", err)
|
||||
}
|
||||
x, err := GMRES(op, b, 2, 50, 1e-12)
|
||||
if err != nil {
|
||||
t.Fatalf("GMRES: %v", err)
|
||||
}
|
||||
for i := range 4 {
|
||||
if math.Abs(x.FloatAt(i)-ref.FloatAt(i)) > 1e-10 {
|
||||
t.Fatalf("x[%d] = %.16g, want %.16g", i, x.FloatAt(i), ref.FloatAt(i))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGMRESErrors(t *testing.T) {
|
||||
b := mustFloats(t, []float64{1, 2}, 2)
|
||||
if _, err := GMRES(func(*core.Array) (*core.Array, error) {
|
||||
return nil, base.Errf("operator failed")
|
||||
}, b, 0, 0, 0); err == nil {
|
||||
t.Fatal("operator error: want an error")
|
||||
}
|
||||
if x, err := GMRES(func(*core.Array) (*core.Array, error) {
|
||||
return nil, base.Errf("operator failed")
|
||||
}, mustFloats(t, []float64{}, 0), 0, 0, 0); err == nil {
|
||||
t.Fatalf("empty system: want an error, got %v", x)
|
||||
}
|
||||
}
|
||||
|
||||
// TestGMRESUnconvergedErrors pins the contract shared with SpSolve and
|
||||
// FindRootSystem: running out of cycles with the tolerance unmet is an
|
||||
// error naming the residual, never a silent approximation.
|
||||
func TestGMRESUnconvergedErrors(t *testing.T) {
|
||||
// A = diag(1, 0.5): a legitimate but slow system for restart-1
|
||||
// GMRES, which crawls towards the solution and cannot reach a
|
||||
// 1e-12 relative residual in 2 cycles.
|
||||
diag := []float64{1, 0.5}
|
||||
op := func(v *core.Array) (*core.Array, error) {
|
||||
out := core.New(core.Float, 2)
|
||||
for i := range 2 {
|
||||
out.RawFloats()[i] = diag[i] * v.FloatAt(i)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
if _, err := GMRES(op, mustFloats(t, []float64{1, 1}, 2), 1, 2, 1e-12); err == nil {
|
||||
t.Fatal("unconverged GMRES: want an error")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNorm2F64LargeScale pins the overflow-safe norm: entries near
|
||||
// 1e154 must not square their way to +Inf.
|
||||
func TestNorm2F64LargeScale(t *testing.T) {
|
||||
if got := norm2F64([]float64{1e154, 1e154}); math.IsInf(got, 1) || math.Abs(got-math.Sqrt2*1e154) > 1e-12*math.Sqrt2*1e154 {
|
||||
t.Fatalf("norm2F64(1e154, 1e154) = %v, want ≈ %g", got, math.Sqrt2*1e154)
|
||||
}
|
||||
if got := norm2F64([]float64{0, 0}); got != 0 {
|
||||
t.Fatalf("norm2F64(zeros) = %v, want 0", got)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user