// Copyright (c) 2026 Petr Balvín (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") } }