52 lines
1.6 KiB
Go
52 lines
1.6 KiB
Go
package optim
|
|||
|
|
|
||
|
|
import (
|
||
|
|
"math"
|
||
|
|
"testing"
|
||
|
|
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
|
|
)
|
||
|
|
|
||
|
|
// A residual linear in the parameters makes one Gauss-Newton step the
|
||
|
|
// exact answer: the normal equations solve to the least-squares
|
||
|
|
// solution in a single iteration, so the fit's first step must land on
|
||
|
|
// it. Every off-diagonal entry of the three-parameter normal matrix is
|
||
|
|
// load-bearing there; a step computed from an asymmetric matrix lands
|
||
|
|
// somewhere else.
|
||
|
|
func TestLevenbergMarquardtSingleStepExactQuadratic(t *testing.T) {
|
||
|
|
// A is 12x3, well conditioned, nothing symmetric in its columns.
|
||
|
|
A := [][]float64{
|
||
|
|
{1, 0.5, -0.25}, {2, -1, 0.75}, {-0.5, 1.5, 2}, {0.25, -0.75, 1},
|
||
|
|
{1.5, 2, -1}, {-2, 0.25, 0.5}, {0.75, -2, 1.5}, {1, 1, 1},
|
||
|
|
{-1, 0.5, 2}, {0.5, 2, -0.5}, {2, 1, -2}, {-0.25, -1.5, 0.25},
|
||
|
|
}
|
||
|
|
truth := []float64{1.5, -0.75, 2.25}
|
||
|
|
y := make([]float64, len(A))
|
||
|
|
for r := range A {
|
||
|
|
y[r] = A[r][0]*truth[0] + A[r][1]*truth[1] + A[r][2]*truth[2]
|
||
|
|
}
|
||
|
|
residual := func(p *core.Array) (*core.Array, error) {
|
||
|
|
pv := p.RawFloats()
|
||
|
|
out := make([]float64, len(A))
|
||
|
|
for r := range A {
|
||
|
|
out[r] = A[r][0]*pv[0] + A[r][1]*pv[1] + A[r][2]*pv[2] - y[r]
|
||
|
|
}
|
||
|
|
return core.FromFloats(out, len(A))
|
||
|
|
}
|
||
|
|
p0, err := core.FromFloats([]float64{0, 0, 0}, 3)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
opts := LMOptions{MaxIterations: 8}
|
||
|
|
pFit, _, err := LevenbergMarquardt(residual, p0, opts)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
got := pFit.RawFloats()
|
||
|
|
for k := range truth {
|
||
|
|
if math.Abs(got[k]-truth[k]) > 1e-9 {
|
||
|
|
t.Fatalf("parameter %d: the fit answered %g, the exact single step is %g", k, got[k], truth[k])
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|