110 lines
3.7 KiB
Go
110 lines
3.7 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package optim
|
|
|
|
import (
|
|
"math"
|
|
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|
"testing"
|
|
)
|
|
|
|
// TestLBFGSRosenbrock checks convergence on the Rosenbrock valley,
|
|
// the standard test for quasi-Newton methods whose narrow curved
|
|
// valley defeats naive gradient descent.
|
|
func TestLBFGSRosenbrock(t *testing.T) {
|
|
rosenbrock := func(p *core.Array) (float64, error) {
|
|
x, y := p.FloatAt(0), p.FloatAt(1)
|
|
return (1-x)*(1-x) + 100*(y-x*x)*(y-x*x), nil
|
|
}
|
|
gradFn := func(p *core.Array) (*core.Array, error) {
|
|
x, y := p.FloatAt(0), p.FloatAt(1)
|
|
out := core.New(core.Float, 2)
|
|
out.RawFloats()[0] = -2*(1-x) - 400*x*(y-x*x)
|
|
out.RawFloats()[1] = 200 * (y - x*x)
|
|
return out, nil
|
|
}
|
|
start, _ := core.FromFloats([]float64{-1.2, 1}, 2)
|
|
point, value, err := MinimiseLBFGS(rosenbrock, gradFn, start, LBFGSOptions{})
|
|
if err != nil {
|
|
t.Fatalf("MinimiseLBFGS: %v", err)
|
|
}
|
|
if math.Abs(point.FloatAt(0)-1) > 1e-5 || math.Abs(point.FloatAt(1)-1) > 1e-5 {
|
|
t.Fatalf("minimiser = (%.10g, %.10g), want (1, 1)", point.FloatAt(0), point.FloatAt(1))
|
|
}
|
|
if value > 1e-10 {
|
|
t.Fatalf("minimum value = %v, want ≈ 0", value)
|
|
}
|
|
}
|
|
|
|
// TestLBFGSNumericGradient checks that finite-difference gradients
|
|
// converge on the same problem without a caller-supplied derivative.
|
|
func TestLBFGSNumericGradient(t *testing.T) {
|
|
f := func(p *core.Array) (float64, error) {
|
|
x, y := p.FloatAt(0), p.FloatAt(1)
|
|
return (1-x)*(1-x) + 100*(y-x*x)*(y-x*x), nil
|
|
}
|
|
start, _ := core.FromFloats([]float64{-1.2, 1}, 2)
|
|
point, _, err := MinimiseLBFGS(f, nil, start, LBFGSOptions{Tolerance: 1e-6})
|
|
if err != nil {
|
|
t.Fatalf("MinimiseLBFGS: %v", err)
|
|
}
|
|
if math.Abs(point.FloatAt(0)-1) > 1e-3 || math.Abs(point.FloatAt(1)-1) > 1e-3 {
|
|
t.Fatalf("minimiser = (%.10g, %.10g), want (1, 1)", point.FloatAt(0), point.FloatAt(1))
|
|
}
|
|
}
|
|
|
|
// TestLBFGSConvergenceSpeed checks that L-BFGS converges in
|
|
// substantially fewer iterations than gradient descent on the same
|
|
// ill-conditioned problem.
|
|
func TestLBFGSConvergenceSpeed(t *testing.T) {
|
|
const n = 10
|
|
// Diagonal quadratic with eigenvalues 1..n.
|
|
f := func(p *core.Array) (float64, error) {
|
|
s := 0.0
|
|
for i := range p.Len() {
|
|
d := p.FloatAt(i) - float64(i+1)
|
|
s += float64(i+1) * d * d
|
|
}
|
|
return s, nil
|
|
}
|
|
gradFn := func(p *core.Array) (*core.Array, error) {
|
|
out := core.New(core.Float, n)
|
|
for i := range n {
|
|
out.RawFloats()[i] = 2 * float64(i+1) * (p.FloatAt(i) - float64(i+1))
|
|
}
|
|
return out, nil
|
|
}
|
|
start, _ := core.FromFloats(make([]float64, n), n)
|
|
point, _, err := MinimiseLBFGS(f, gradFn, start, LBFGSOptions{})
|
|
if err != nil {
|
|
t.Fatalf("MinimiseLBFGS: %v", err)
|
|
}
|
|
for i := range n {
|
|
want := float64(i + 1)
|
|
if math.Abs(point.FloatAt(i)-want) > 1e-6 {
|
|
t.Fatalf("x[%d] = %.10g, want %.10g", i, point.FloatAt(i), want)
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestLBFGSInvalid pins the error contract.
|
|
func TestLBFGSInvalid(t *testing.T) {
|
|
f := func(p *core.Array) (float64, error) { return 0, nil }
|
|
empty, _ := core.FromFloats(nil, 0)
|
|
if _, _, err := MinimiseLBFGS(f, nil, empty, LBFGSOptions{}); err == nil {
|
|
t.Fatal("expected an error for an empty starting point")
|
|
}
|
|
cx, _ := core.FromComplexes([]complex128{1}, 1)
|
|
if _, _, err := MinimiseLBFGS(f, nil, cx, LBFGSOptions{}); err == nil {
|
|
t.Fatal("expected an error for a complex starting point")
|
|
}
|
|
// Objective that errors must propagate.
|
|
failing := func(p *core.Array) (float64, error) { return 0, base.Errf("objective exploded") }
|
|
start, _ := core.FromFloats([]float64{1}, 1)
|
|
if _, _, err := MinimiseLBFGS(failing, nil, start, LBFGSOptions{}); err == nil {
|
|
t.Fatal("expected the objective's error to propagate")
|
|
}
|
|
}
|