Files
tensor/examples/ode-fit/main.go
T

136 lines
3.7 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command ode-fit fits the parameters of a damped oscillator to
// endpoint measurements by the adjoint method: AdjointODE hands back
// dL/dθ for every parameter at the cost of one extra solve, and plain
// gradient descent walks the parameters to the truth. This is the
// data-assimilation loop no other Go library expresses.
//
// Usage: go run ./examples/ode-fit
package main
import (
"fmt"
"log"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
// coreDynamics is y' = (v, −c·v − ω²·y) over plain arrays, the shape
// IntegrateODE wants; it generates the data.
func coreDynamics(c, omega float64) func(float64, *tensor.Array) (*tensor.Array, error) {
return func(t float64, y *tensor.Array) (*tensor.Array, error) {
return tensor.FromFloats([]float64{
y.FloatAt(1),
-c*y.FloatAt(1) - omega*omega*y.FloatAt(0),
}, 2)
}
}
// graphDynamics is the same equation with (c, ω) as differentiable
// leaves, the shape AdjointODE wants.
func graphDynamics(c, omega *tensor.Tensor) func(float64, *tensor.Tensor) (*tensor.Tensor, error) {
return func(t float64, y *tensor.Tensor) (*tensor.Tensor, error) {
yv, err := y.Slice(0, 0, 1)
if err != nil {
return nil, err
}
vv, err := y.Slice(0, 1, 2)
if err != nil {
return nil, err
}
w2, err := omega.Pow(2)
if err != nil {
return nil, err
}
damping, err := c.Mul(vv)
if err != nil {
return nil, err
}
restoring, err := w2.Mul(yv)
if err != nil {
return nil, err
}
acc, err := damping.Add(restoring)
if err != nil {
return nil, err
}
neg, err := acc.Scale(-1)
if err != nil {
return nil, err
}
return vv.Concat(neg, 0)
}
}
func main() {
const (
trueC = 0.8
trueOmega = 3.0
)
y0, err := tensor.FromFloats([]float64{1, 0}, 2)
if err != nil {
log.Fatal(err)
}
// Measurements of y(t) at three times from the true system.
times := []float64{0.4, 0.8, 1.2, 1.6, 2.0, 2.4}
data := make([]float64, len(times))
trueF := coreDynamics(trueC, trueOmega)
for i, T := range times {
end, err := tensor.IntegrateODE(trueF, 0, T, y0, tensor.ODEOptions{})
if err != nil {
log.Fatal(err)
}
data[i] = end.FloatAt(0)
}
// The fit: gradient descent on L = Σ (y(Tᵢ; θ) − dataᵢ)² with the
// gradient from one adjoint pass per data point.
cVal, wVal := 0.3, 1.8
const rate = 0.04
for iter := 1; iter <= 400; iter++ {
gc, gw := 0.0, 0.0
loss := 0.0
forwardF := coreDynamics(cVal, wVal)
cArr, _ := tensor.FromFloats([]float64{cVal}, 1)
wArr, _ := tensor.FromFloats([]float64{wVal}, 1)
c := tensor.FromArray(cArr, true)
omega := tensor.FromArray(wArr, true)
adjF := graphDynamics(c, omega)
for i, T := range times {
end, err := tensor.IntegrateODE(forwardF, 0, T, y0, tensor.ODEOptions{})
if err != nil {
log.Fatal(err)
}
res := end.FloatAt(0) - data[i]
loss += res * res
// dL/dy(T) = 2·res on the position component only; the
// velocity component carries no loss.
seed, err := tensor.FromFloats([]float64{2 * res, 0}, 2)
if err != nil {
log.Fatal(err)
}
_, paramGrads, err := tensor.AdjointODE(adjF, []*tensor.Tensor{c, omega},
0, T, y0, seed, tensor.ODEOptions{})
if err != nil {
log.Fatal(err)
}
gc += paramGrads[0].FloatAt(0)
gw += paramGrads[1].FloatAt(0)
}
if iter%100 == 0 {
fmt.Printf("iter %3d c = %.4f omega = %.4f loss = %.3e\n", iter, cVal, wVal, loss)
}
cVal -= rate * gc
wVal -= rate * gw
}
fmt.Printf("fitted c = %.4f (true %.4f), omega = %.4f (true %.4f)\n",
cVal, trueC, wVal, trueOmega)
if math.Abs(cVal-trueC) > 0.05 || math.Abs(wVal-trueOmega) > 0.05 {
log.Fatal("the fit did not converge to the truth")
}
}