Files
tensor/integrate/wave_bench_test.go
T

153 lines
4.1 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
package integrate
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Benchmarks for the library-side cost of the solver drivers: the
// same workloads the package benchmarks run, with the fixture's own
// per-call allocation removed, so the allocation count a run reports
// is the library's alone. Every fixture derivative refills one output
// array the driver copies out of immediately, which the f contract
// allows: no solver retains the returned array.
// wavePooledDecay builds the diagonal stiff system y' = rate·(i+1)·y_i
// with an f that refills one shared output array per call.
func wavePooledDecay(n int, rate float64) func(t float64, y *core.Array) (*core.Array, error) {
rates := make([]float64, n)
for i := range rates {
rates[i] = rate * float64(i+1)
}
out := core.New(core.Float, n)
vals := out.RawFloats()
return func(t float64, y *core.Array) (*core.Array, error) {
ys := y.RawFloats()
for i := range n {
vals[i] = rates[i] * ys[i]
}
return out, nil
}
}
// waveVector wraps a fixed literal as a rank-1 array.
func waveVector(b *testing.B, vals []float64) *core.Array {
b.Helper()
a, err := core.FromFloats(vals, len(vals))
if err != nil {
b.Fatal(err)
}
return a
}
// waveOnes returns a vector of n ones.
func waveOnes(n int) []float64 {
vals := make([]float64, n)
for i := range vals {
vals[i] = 1
}
return vals
}
// BenchmarkWaveLibraryRK4 runs the fixed-step RK4 loop over the linear
// system of BenchmarkIntegrateRK4, with the derivative evaluation
// allocation-free: the reported allocations are the driver's own.
func BenchmarkWaveLibraryRK4(b *testing.B) {
const n = 16
f := wavePooledDecay(n, -0.25)
start := waveVector(b, waveOnes(n))
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateRK4(f, 0, 10, start, 2000); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkWaveLibraryBDF2 runs the variable-step BDF2 loop over the
// stiff diagonal system of BenchmarkIntegrateBDF2, allocation-free on
// the fixture side.
func BenchmarkWaveLibraryBDF2(b *testing.B) {
const n = 4
f := wavePooledDecay(n, -100)
start := waveVector(b, waveOnes(n))
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateBDF2(f, 0, 1, start, ODEOptions{}); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkWaveLibraryBDFVar runs the variable-order BDF loop over the
// stiff diagonal system of BenchmarkBDFVarStiff, allocation-free on the
// fixture side.
func BenchmarkWaveLibraryBDFVar(b *testing.B) {
const n = 32
f := wavePooledDecay(n, -100)
start := waveVector(b, waveOnes(n))
opts := BDFVarOptions{RelTol: 1e-6, AbsTol: 1e-9}
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateBDFVar(f, 0, 1, start, opts); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkWaveLibraryROS4 runs the ROS4 loop over the stiff diagonal
// system of BenchmarkROS4Stiff, allocation-free on the fixture side.
func BenchmarkWaveLibraryROS4(b *testing.B) {
const n = 32
f := wavePooledDecay(n, -100)
start := waveVector(b, waveOnes(n))
opts := ODEOptions{RelTol: 1e-6, AbsTol: 1e-9}
b.ReportAllocs()
for b.Loop() {
if _, err := IntegrateROS4(f, 0, 1, start, opts); err != nil {
b.Fatal(err)
}
}
}
// BenchmarkWaveLibraryVerlet runs the velocity Verlet loop over the
// coupled oscillators of BenchmarkIntegrateVerlet, with the
// acceleration evaluation allocation-free: the reported allocations
// are the driver's own.
func BenchmarkWaveLibraryVerlet(b *testing.B) {
n := 32
q0 := waveOnes(n)
for i := range q0 {
q0[i] = math.Sin(float64(i))
}
out := core.New(core.Float, n)
vals := out.RawFloats()
accel := func(q *core.Array) (*core.Array, error) {
qs := q.RawFloats()
for i := range n {
l, r := 0.0, 0.0
if i > 0 {
l = qs[i-1]
}
if i < n-1 {
r = qs[i+1]
}
vals[i] = l - 2*qs[i] + r
}
return out, nil
}
qs := waveVector(b, q0)
ps := waveVector(b, waveOnes(n))
b.ReportAllocs()
for b.Loop() {
if _, _, err := IntegrateVerlet(accel, 0, 10, qs, ps, 500); err != nil {
b.Fatal(err)
}
}
}