136 lines
3.7 KiB
Go
136 lines
3.7 KiB
Go
// 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")
|
|||
|
|
}
|
|||
|
|
}
|